ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

JAX Export 形状多态(Shape Polymorphism)实战指南:一次导出、多形状复用

2026/9/20 23:52:35 拓冰建站 浏览量
JAX Export 形状多态(Shape Polymorphism)实战指南:一次导出、多形状复用 机器学习深度学习【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址https://gitcode.com/gh_mirrors/jax/jax点击查看免费下载形状多态shape polymorphism是 JAX 导出jax.export体系中的一项关键能力当函数以 JIT 模式运行时JAX 会针对每种输入类型与形状组合重新追踪、降级到 StableHLO 并编译而一旦函数被导出并在另一系统上反序列化Python 源码已不可用无法再次追踪与降级。形状多态允许导出的函数服务一整族输入形状——函数在导出时只追踪、降级一次Exported对象中携带足以在多种具体形状上编译执行的信息。本文基于 docs/export/shape_poly.md 官方文档并结合仓库源码jax/_src/export/shape_poly.py、jax/export.py深入讲解其用法、正确性语义与常见错误排查读完即可在项目中使用符号维度导出形状无关的模型与算子。为什么需要形状多态常规的 JIT 工作流中函数f的追踪与编译是按形状缓存的import jax from jax import export from jax import numpy as jnp导出之后情况则不同Exported对象没有 Python 源码无法对新的输入形状重新追踪。因此 JAX 提供形状多态使导出的程序能在一个符号形状symbolic shapes的族上编译执行。实现方式是在导出时指定包含维度变量的形状def f(x): # f: f32[a, b] return jnp.concatenate([x, x], axis1) # 构造符号维度变量 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)需要特别强调的是此类函数在每次以具体输入形状调用时仍然会按需重新编译被保存下来的只是追踪tracing与降级lowering的结果。从源码看symbolic_shape的完整签名见 jax/_src/export/shape_poly.py为def symbolic_shape(shape_spec: str | None, *, constraints: Sequence[str] (), scope: SymbolicScope | None None, like: Sequence[int | None] | None None, ) - Sequence[DimSize]:shape_spec形状规格字符串None等价于...。它是元组的字符串表示括号可省略逗号分隔各个维度表达式维度表达式可以是整数常量、维度变量字母开头的字母数字串、e1 e2、e1 - e2、e1 * e2、floordiv(e1, e2)、mod(e1, e2)、max(e1, e2)、min(e1, e2)constraints对符号维度表达式的约束序列形式为e1 e2、e1 e2或e1 e2scope可选的 jax.export.SymbolicScope 对象指定后表达式创建于该作用域scope与constraints不能同时指定同时指定会抛出ValueErrorlike当shape_spec含有占位符_、...时用该形状填充占位符被用于填充的like维度不能是None。返回的维度表达式对象为_DimExpr见 jax/_src/export/shape_poly.py它重载了绝大多数整数运算符因此大多数情况下可以像整数常量一样使用详见下文与维度变量计算一节。用 symbolic_args_specs 批量生成参数规格除手工构造jax.ShapeDtypeStruct外jax.export.symbolic_args_specs可以基于多态形状规格直接构造一个jax.ShapeDtypeStruct的 pytreeimport numpy as np 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, ...中的占位符...会根据具体参数(x, y)的形状填充...代表 0 个或多个维度而_只代表一个维度。symbolic_args_specs支持参数的 pytree——它从参数中获取 pytree 结构、dtype 并填充占位符最终构造一个与参数结构匹配的jax.ShapeDtypeStructpytree。当一条规格要应用于多个参数时shapes_specs可以是一个 pytree 前缀。关于可选参数与 pytree 前缀的匹配规则可参考仓库中的 docs/working-with-pytrees.md。其源码签名见 jax/_src/export/shape_poly.py中shapes_specs的取值有三种None所有参数均为静态形状单个字符串作为shape_spec应用于所有参数与args前缀匹配的 pytree对不同参数分别指定。内部实现通过tree_util.broadcast_prefix把规格广播到参数树再从参数提取形状与 dtypeshape_and_dtype_jax_array完成占位符填充。形状规格写法示例((b, _, _), None)适用于两个参数的函数第一个参数是 3D 数组批次维b为符号维度其余维度以及第二个参数的形状按实际参数特化。同样的规格也适用于第一个参数是前导维度相同、尾随维度可不同的 3D 数组 pytree 的情况None表示第二个参数不做符号化等价于...。((batch, ...), (batch,))指定两个参数前导维度一致第一个参数秩至少为 1第二个参数秩为 1。形状多态的正确性契约形状多态的正确性需要被严格定义。对于任意 JAX 函数f、包含符号形状的参数规格arg_spec以及形状与arg_spec匹配的具体参数arg若 JAX 原生执行成功res f(arg)且符号形状导出成功exp export.export(f)(arg_spec)则编译并运行导出对象必然成功且结果一致res exp.call(arg)。关键在于理解f(arg)有权重新调用 JAX 的追踪机制实际上它对每种不同的具体arg形状都会这样做而exp.call(arg)的执行不能再使用 JAX 追踪——它可能运行在源码不可用的环境中。要保证这种正确性并不容易在最困难的情形下导出会直接失败本章剩余内容就是讲解如何应对这些失败。与维度变量计算JAX 会跟踪所有中间结果的形状。当形状依赖维度变量时JAX 会把它们计算为符号维度表达式。维度变量代表大于等于 1 的整数值。符号表达式可以表示对维度表达式与整数int、np.int或任何可用operator.index转换的对象应用算术运算符加、减、乘、整除floordiv、取模mod包括 NumPy 变体np.sum、np.prod等的结果。这些符号维度可以用于 JAX 原语与 API 的形状参数例如jnp.reshape、jnp.arange、切片索引等。例如展平 2D 数组时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]),)维度表达式与 JAX 数组的互相转换可以用jnp.array(x.shape[0])甚至jnp.array(x.shape)把维度表达式显式转换为 JAX 数组。结果可以像普通 JAX 数组一样使用但不能再用作形状中的维度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)) # 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)].与非常量标量混合时自动转换为数组当符号维度与非整数如float、np.float、np.ndarray或 JAX 数组参与运算时会自动通过jnp.array转换为 JAX 数组。下面函数中所有x.shape[0]的出现都会被隐式转换为jnp.array(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))一个典型场景是计算平均值注意x.shape[0]如何自动变成 JAX 数组exp export.export(jax.jit( lambda x: jnp.sum(x, axis0) / x.shape[0]))( jax.ShapeDtypeStruct(export.symbolic_shape(b, c), jnp.int32)) exp.call(jnp.arange(12, dtypejnp.int32).reshape((3, 4))) # Array([4., 5., 6., 7.], dtypefloat32)形状多态下的常见错误大多数 JAX 代码假设数组形状是整数元组而形状多态下某些维度可能是符号表达式这会导致多种错误。常规形状检查错误首先是常见的 JAX 形状检查错误v, export.symbolic_shape(v,) export.export(jax.jit(lambda x, y: x y))( jax.ShapeDtypeStruct((v,), dtypenp.int32), jax.ShapeDtypeStruct((4,), dtypenp.int32)) # Traceback (most recent call last): # TypeError: add got incompatible shapes for broadcasting: (v,), (4,). export.export(jax.jit(lambda x: jnp.matmul(x, x)))( jax.ShapeDtypeStruct((v, 4), dtypenp.int32)) # Traceback (most recent call last): # TypeError: dot_general requires contracting dimensions to have the same shape, got (4,) and (v,).修复上面 matmul 例子只需把参数形状指定为(v, v)。符号维度比较部分支持JAX 内部有大量涉及形状的相等与不等比较用于形状检查甚至用于为某些原语选择实现比较规则如下相等若两个符号维度在所有维度变量的取值下必然相同则相等判定为True如b b 2*b否则判定为False。该行为的严重后果见下文等值比较的注意点不相等恒为相等的否定不等式部分支持与部分相等类似但会考虑维度变量取严格正整数这一事实。例如b 1、b 0、2 * a b 3为True而b 2、a b、a - b 0无法判定会抛出异常。当比较无法归结为布尔值时会抛出InconclusiveDimensionOperation该类定义于 jax/_src/export/shape_poly.py继承自core.InconclusiveDimensionOperationimport 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)) # 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.遇到InconclusiveDimensionOperation可以尝试以下策略若代码使用了内置的max/min或np.max/np.min改用core.max_dim与core.min_dim——它们把不等式比较延迟到编译期此时形状已确定。这两个函数定义于 jax/_src/core.py语义为同时适用于常量维度与符号维度的max/min尝试用core.max_dim/core.min_dim重写条件式例如把d if d 0 else 0改写为core.max_dim(d, 0)改写代码减少对维度必须是整数的依赖利用符号维度在多数算术运算中可作为整数鸭子类型的事实例如把int(d) 5写成d 5指定符号约束见下文。用户指定的符号约束默认情况下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 需要验证切片大小不超过轴大小_ export.export(jax.jit(lambda x: x[0:16]))( jax.ShapeDtypeStruct(export.symbolic_shape(b 15), dtypenp.int32))这类隐式约束参与比较判定并在编译期检查见下文形状断言错误。还可以指定显式约束# 引入带约束的维度变量 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 对符号约束的推理支持有限变量与常量比较的约束收益最大例如由a 16与b 8可以推断a 2*b 32涉及复杂表达式的约束推理能力有限例如从a b 8可推断a - b 8但不能推断a 9未来可能改进相等约束被当作归一化规则例如floordiv(a, b) c会把所有左侧出现替换为右侧。相等约束的左侧必须是因子的乘积如a * b、4 * a、floordiv(a, b)顶层不能包含加法或减法。符号约束还能绕过 JAX 推理机制的局限。例如下面的代码中JAX 需要证明切片大小x.shape[0] % 3即符号表达式mod(b, 3)不超过轴大小b。这对所有严格正数b都成立但 JAX 的符号比较规则证明不了因此报错from jax import lax b, export.symbolic_shape(b) f lambda x: lax.slice_in_dim(x, 0, x.shape[0] % 3) export.export(jax.jit(f))( jax.ShapeDtypeStruct((b,), dtypenp.int32)) # Traceback (most recent call last): # jax._src.export.shape_poly.InconclusiveDimensionOperation: Symbolic dimension comparison b mod(b, 3) 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.一种解决方式是把代码限制在轴大小为 3 的倍数的场景把b换成3*b此时 JAX 可以把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))与隐式约束一样显式符号约束也在编译期使用同一机制检查。从源码看SymbolicScopejax/_src/export/shape_poly.py负责持有约束它解析并存储显式约束_explicit_constraints缓存_bounds_cache并把相等约束转化为_normalization_rules归一化规则约束字符串必须包含、或之一无法满足的约束如常量差为负的会直接抛出ValueError。符号维度作用域符号约束存储在 jax.export.SymbolicScope 对象中每次调用symbolic_shape都会隐式创建一个作用域。不要把不同作用域的符号表达式混用。例如下面的代码会失败因为a1与a2来自不同作用域a1, export.symbolic_shape(a,) a2, export.symbolic_shape(a,, constraints(a 8,)) a1 a2 # 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单次symbolic_shape调用产生的表达式共享同一作用域可以参与算术运算运算结果也共享该作用域。作用域可以复用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 cJAX 追踪使用以形状为键的部分缓存因此打印相同的符号形状若作用域不同仍被视为不同。等值比较的注意点相等比较对b 1 b或b 0返回False此时可确定所有取值下维度都不同但对b 1和a b也返回False——这是不健全unsound的按道理应当抛出core.InconclusiveDimensionOperation因为在某些取值下应为True、另一些取值下应为False。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则不健全。维度变量必须能从输入形状求解目前调用导出对象时传入维度变量值的唯一途径是通过数组参数形状间接推断。例如b的值可以在调用点从第一个参数的类型f32[b]推断。这对大多数用例都适用也与 JIT 函数的调用约定一致。但有时想导出一个由整数值参数化的函数该整数值决定程序中的若干形状。例如导出一个由k参数化、结果形状取决于k的my_top_kdef 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) # Traceback (most recent call last): # KeyError: 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 当前无法求解a, export.symbolic_shape(a) export.export(jax.jit(lambda x: x.shape[0]))( jax.ShapeDtypeStruct((a * a,), dtypenp.int32)) # 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].形状断言错误JAX 假设维度变量取值于严格正整数该假设在针对具体输入形状编译时会得到检查。例如给定符号输入形状(b, b, 2*d)JAX 在接收实际参数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)) # 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), .这些错误发生在编译前的预处理步骤中。符号维度的除法部分支持JAX 会尝试化简除法与取模运算例如(a * b a) // (b 1) a、6*a 4 % 3 1。具体来说JAX 处理两种情况(a) 无余数(b) 除数是常数此时可能存在常数余数。例如下面的代码在推断reshape的目标维度时发生除法错误b, export.symbolic_shape(b) export.export(jax.jit(lambda x: x.reshape((2, -1))))( jax.ShapeDtypeStruct((b,), dtypenp.int32)) # Traceback (most recent call last): # jax._src.core.InconclusiveDimensionOperation: Cannot divide evenly the sizes of shapes (b,) and (2, -1). # The remainder mod(b, - 2) should be 0.而下面这些写法可以成功b, export.symbolic_shape(b) # 指定第一个维度是 4 的倍数 exp export.export(jax.jit(lambda x: x.reshape((2, -1))))( jax.ShapeDtypeStruct((4*b,), dtypenp.int32)) exp.out_avals # (ShapedArray(int32[2,2*b]),) # 指定某个维度为偶数 exp export.export(jax.jit(lambda x: x.reshape((2, -1))))( jax.ShapeDtypeStruct((b, 5, 6), dtypenp.int32)) exp.out_avals # (ShapedArray(int32[2,15*b]),)调试形状多态问题调试时首先参考导出体系的整体调试文档 docs/export/export.md。此外可以对形状细化shape refinement过程单独调试——该过程在编译期对含维度变量或多平台支持的模块执行。若形状细化期间出错可以设置JAX_DUMP_IR_TO环境变量转储形状细化之前的 HLO 模块文件名为..._before_refine_polymorphic_shapes.mlir该模块应已具有静态输入形状# 转储 IR 并开启 C 侧形状细化日志OSS 版 JAX_DUMP_IR_TO/tmp/export.dumps/ TF_CPP_VMODULErefine_polymorphic_shapes3 python tests/shape_poly_test.py ShapePolyTest.test_simple_unary -v3JAX_DUMP_IR_TO指向转储目录TF_CPP_VMODULErefine_polymorphic_shapes3在 OSS 下开启refine_polymorphic_shapes各阶段的日志Google 内部对应--vmodulerefine_polymorphic_shapes3示例命令同时运行了仓库中的形状多态测试用例 tests/shape_poly_test.pyShapePolyTest.test_simple_unary可作为复现与验证的最小入口。小结形状多态是 JAX 导出体系中连接一次导出与多种形状复用的核心机制。掌握symbolic_shape与symbolic_args_specs的用法、理解维度变量的算术与比较语义、善用隐式/显式符号约束、牢记维度变量必须能从输入形状求解以及作用域不能混用等约束就能为模型与算子写出形状无关的导出工件。遇到InconclusiveDimensionOperation或形状断言错误时优先考虑用core.max_dim/core.min_dim延迟比较、改写依赖整数语义的代码、调整形状规格如4*b、b 15或补充显式约束必要时借助JAX_DUMP_IR_TO与TF_CPP_VMODULE定位形状细化阶段的问题。相关源码与测试可进一步阅读 jax/_src/export/shape_poly.py、jax/_src/core.py 与 tests/shape_poly_test.py。赞分享机器学习深度学习【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址https://gitcode.com/gh_mirrors/jax/jax点击查看免费下载相关推荐JAX Shape Polymorphism 完全指南用符号维度实现一次导出、多形状复用JAX Shape Polymorphism 完全指南用符号维度实现一次导出、多形状复用 导读 本文围绕 JAX 官方文档 docs/501/shape po人工智能机器学习深度学习编译器高性能计算JAX Shape Polymorphism 实战指南用符号维度导出可复用的多形状计算JAX Shape Polymorphism 实战指南用符号维度导出可复用的多形状计算 本篇指南系统讲解 JAX 中 jax.export 的 Shape P人工智能机器学习深度学习编译器高性能计算JAX 导出与序列化实战指南从 StableHLO 导出到形状多态与跨平台部署JAX 导出与序列化实战指南从 StableHLO 导出到形状多态与跨平台部署 本指南以 docs/export/index.rst https://link机器学习深度学习上一篇highlight.php显式模式vs自动检测模式哪种更适合你的项目下一篇BaRMIe代码架构解析揭秘RMI代理与Payload注入机制创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考