
JAX Shape Polymorphism 完全指南用符号维度实现一次导出、多形状复用【免费下载链接】jaxComposable transformations of PythonNumPy 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], axis1)问题在于当函数被export导出、序列化并在另一台机器上反序列化之后Python 源码已经不在现场无法重新追踪re-trace和重新降级re-lower。此时若输入形状发生变化例如 batch 维度从 8 变成 16传统方案就会失效。Shape polymorphism 正是为这一场景而生在导出阶段用符号形状symbolic shapes进行追踪与降级把维度写成一个或多个维度变量dimension variables而Exported对象中保存了足以在多种具体输入形状下编译执行的全部信息。函数在调用时仍会按需重新编译但只有编译这一步发生在调用现场追踪与降级只发生一次、随导出而固化。这一点在官方文档中明确强调也是理解该特性性能与语义的关键。2. 核心 API 一jax.export.symbolic_shape与symbolic_args_specs2.1 用symbolic_shape构造符号形状jax.export.symbolic_shape(shape_spec, *, constraints(), scopeNone, likeNone)接受一个字符串形式的形状规格返回维度表达式对象的元组类型为_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]),) # 之后可以用具体形状调用这里 a3, b4无需重新追踪 f。 res exp.call(np.ones((3, 4), dtypenp.int32)) res.shape (3, 8)观察out_avals中的int32[a,2*b]jnp.concatenate([x, x], axis1)的输出形状被 JAX 自动计算为2*b这一符号维度表达式。维度表达式对象重载了绝大多数整数运算符因此在大多数场景下可以像使用整数常量一样参与算术、切片与形状计算详见 docs/501/shape-polymorphism.md 与第 4 节。2.2 用symbolic_args_specs从真实参数构造规格 pytreejax.export.symbolic_args_specs(args, shapes_specs, *, constraints(), scopeNone)的用途是基于实际具体形状参数构造出与之一一对应的jax.ShapeDtypeStructpytree其中被占位符覆盖的维度替换为符号维度dtype 则沿用真实参数。看文档中的完整示例 def f1(x, y): # x: f32[a, 1], y : f32[a, 4] ... return x y # 假设你已有具体形状的真实参数 x np.ones((3, 1), dtypenp.int32) y np.ones((3, 4), dtypenp.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, likes, scopescope)完成占位符填充最终用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, dtypenp.int32)) Array([3, 4, 5], dtypeint32) 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 [TracedShapedArray(int32[], weak_typeTrue)withDynamicJaxprTrace(level1/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, dtypejnp.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, dtypejnp.int32), ... x x.shape[0] jnp.sin(x.shape[0]))))( ... jax.ShapeDtypeStruct(export.symbolic_shape(b), jnp.int32)) exp.out_avals (ShapedArray(float32[], weak_typeTrue), ShapedArray(int32[5]), ShapedArray(float32[b], weak_typeTrue)) exp.call(jnp.ones((3,), jnp.int32)) (Array(8., dtypefloat32, weak_typeTrue), Array([ 3, 2, 1, 0, -1], dtypeint32), Array([4.14112, 4.14112, 4.14112], dtypefloat32, weak_typeTrue))另一个典型场景是求平均jnp.sum(x, axis0) / x.shape[0]中x.shape[0]同样被自动转成数组参与除法得到正确结果Array([4., 5., 6., 7.], dtypefloat32)。该自动转换机制的底层实现可见 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,), dtypenp.int32), # doctest: IGNORE_EXCEPTION_DETAIL ... jax.ShapeDtypeStruct((4,), dtypenp.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), dtypenp.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. 符号维度的比较部分支持与InconclusiveDimensionOperationJAX 内部存在大量涉及形状的相等/不等比较用于形状检查甚至选择某些原语的实现。符号维度下的比较规则如下相等比较带一个警示若两个符号维度在所有维度变量取值下都相等则结果为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), dtypenp.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。文档示例若不写 15JAX 无法证明切片大小 16 不超过轴大小b导出会失败 _ export.export(jax.jit(lambda x: x[0:16]))( ... jax.ShapeDtypeStruct(export.symbolic_shape(b 15), dtypenp.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), dtypenp.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,), dtypenp.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 的追踪缓存以形状为部分键打印相同的符号形状若来自不同作用域也会被视为不同。复用作用域的两种方式通过scopea.scope复用已有维度变量的作用域此时不能再附加新约束或显式创建SymbolicScope a, export.symbolic_shape(a,, constraints(a 8,)) b, export.symbolic_shape(b,, scopea.scope) # 复用 a 的作用域 a b # 允许 b a my_scope export.SymbolicScope() c, export.symbolic_shape(c, scopemy_scope) d, export.symbolic_shape(d, scopemy_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.pysymbolic_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, dtypenp.int32).reshape((4, 10)) # 用静态 k3 导出。由于 k 出现在形状中必须放进 static_argnums。 exp_static_k export.export(jax.jit(my_top_k, static_argnums0))(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]], dtypeint32) # 现在尝试用符号 k 导出以便导出后再选 k。 k, export.symbolic_shape(k, constraints[k 10]) export.export(jax.jit(my_top_k, static_argnums0))(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), dtypenp.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), dtypenp.int32), x) Array([[ 9, 8, 7], [19, 18, 17], [29, 28, 27], [39, 38, 37]], dtypeint32)另一种报错场景是维度变量虽然出现在输入形状中但以 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,), dtypenp.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), dtypenp.int32)) exp.call(np.ones((3, 3, 5), dtypenp.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_VMODULErefine_polymorphic_shapes3OSS 环境Google 内部用--vmodulerefine_polymorphic_shapes3。文档给出的完整命令示例# Log from python JAX_DUMP_IR_TO/tmp/export.dumps/ TF_CPP_VMODULErefine_polymorphic_shapes3 python tests/shape_poly_test.py ShapePolyTest.test_simple_unary -v3该命令同时演示了如何运行官方测试套件中的用例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, axis0)) / 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), dtypenp.float32)) assert r.shape (n, 3)仓库中可交叉验证的素材包括核心实现jax/_src/export/shape_poly.pysymbolic_shape、symbolic_args_specs、symbolic_dim_bounds、SymbolicScope、shape_assertion_p等导出主流程与Exported.calljax/_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 PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考