ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

如何用符号 shape 导出 JAX 函数,让一个 Exported 对象支持一整族输入形状

如何用符号 shape 导出 JAX 函数,让一个 Exported 对象支持一整族输入形状 如何用符号 shape 导出 JAX 函数让一个 Exported 对象支持一整族输入形状【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jaxJAX 在 JIT 模式下会为每一种输入类型与形状的组合分别 trace、lower 到 StableHLO 并编译。函数导出jax.export之后再在另一台系统上反序列化时Python 源码已经不可用无法重新 trace 和 lower。JAX export 的 shape polymorphism形状多态功能就是为了解决这个问题导出时给部分维度指定 dimension variable符号维度函数只需在导出时 trace、lower 一次得到的Exported对象就能在多个具体输入形状上编译和执行。本文以 docs/export/shape_poly.md 为主线演示如何完成导出、验证Exported对象并处理导出和调用阶段会遇到的典型错误。用符号 shape 导出函数主路径分三步构造符号维度变量、用它构造参数规格、导出并检查。下面的代码块可直接执行import jax import numpy as np from jax import export from jax import numpy as jnp def f(x): # x: int32[a, b] return jnp.concatenate([x, x], axis1) # 1. 构造符号维度变量。字符串 a, b 被解析为维度表达式对象类型 _DimExpr a, b export.symbolic_shape(a, b) # 2. 用符号维度构造 shape代替整数常量 x_shape (a, b) # 3. 用包含符号 shape 的 ShapeDtypeStruct 导出 exp: export.Exported export.export(jax.jit(f))( jax.ShapeDtypeStruct(x_shape, jnp.int32)) print(exp.in_avals) # (ShapedArray(int32[a,b]),) print(exp.out_avals) # (ShapedArray(int32[a,2*b]),)上面两行 avals 输出为文档示例。导出成功的判据就是in_avals/out_avals里出现了符号表达式输入是int32[a,b]输出是int32[a, 2*b]而不是被具体数值钉死的形状。之后可以用任意满足规格的具体形状调用且不会重新 tracef# 具体形状 (3, 4)即 a3, b4 res exp.call(np.ones((3, 4), dtypenp.int32)) print(res.shape) # (3, 8)文档示例这里要清楚一个边界这种导出函数在每个具体输入形状上仍然会按需重新编译被保存的只是 trace 和 lower 的结果。shape polymorphism 省去的是另一台机器上必须重新 trace 和 lower不是省去全部编译。维度表达式重载了大部分整数运算符所以在大多数地方可以像用整数常量一样使用它们。例如把 2D 数组拍平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) print(exp.out_avals) # (ShapedArray(int32[4*b]),)文档示例文档同时给出了正确性保证对任意 JAX 函数f、含符号 shape 的参数规格arg_spec以及形状匹配arg_spec的具体参数arg——若原生执行f(arg)成功、符号导出也成功则exp.call(arg)编译运行成功且结果与f(arg)相同。这是一个 Exported 对象支持一整族输入形状这一承诺的依据。多参数或 pytree 参数symbolic_args_specs当函数有多个参数或参数是 pytree 时手工为每个参数构造ShapeDtypeStruct会很繁琐。export.symbolic_args_specs可以从真实参数出发按多态 shape 规格生成整棵 pytree 的参数规格def f1(x, y): # x: int32[a, 1], y: int32[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) print(exp.in_avals) # (ShapedArray(int32[a,1]), ShapedArray(int32[a,4]))文档示例规格字符串里有两类占位符...代表 0 个或多个维度_代表恰好 1 个维度占位符的具体取值从实际参数args的 shape 和 dtype 填入。规格本身也可以是 pytree 前缀即一份规格应用到多个参数。文档给出的两个示例规格((b, _, _), None)函数有两个参数第一个是 3D 数组且 leading batch 维符号化其余维度按实际参数特化None表示第二个参数不做符号化等价于...。若第一个参数本身是同 leading 维、尾部维度可能不同的 3D 数组 pytree同一份规格也适用。((batch, ...), (batch,))两个参数 leading 维匹配第一个 rank 至少 1第二个 rank 为 1。注意...与_的区别前者可展开为任意数量的维度后者只匹配一个维度。处理导出时的维度运算错误大多数 JAX 代码假定 shape 是整数元组而形状多态下部分维度是符号表达式导出时会出现两类典型错误。普通 shape 检查错误广播和矩阵乘法等在符号维度上不匹配时会直接报TypeErrorv, export.symbolic_shape(v,) # 广播不兼容(v,) 与 (4,) export.export(jax.jit(lambda x, y: x y))( jax.ShapeDtypeStruct((v,), dtypenp.int32), jax.ShapeDtypeStruct((4,), dtypenp.int32)) # TypeError: add got incompatible shapes for broadcasting: (v,), (4,). # matmul 收缩维不一致 export.export(jax.jit(lambda x: jnp.matmul(x, x)))( jax.ShapeDtypeStruct((v, 4), dtypenp.int32)) # TypeError: dot_general requires contracting dimensions # to have the same shape, got (4,) and (v,).以上错误信息为文档示例。文档给出的修复方式把 matmul 示例的参数规格改写成形状(v, v)让收缩维一致。InconclusiveDimensionOperation比较无法判定JAX 内部大量使用 shape 相等/不等比较做 shape 检查、选择原语实现等。这些比较在符号维度下是部分支持的相等比较在两个表达式对所有取值都相等时为True如b b 2*b否则为False不等式只在能依据维度变量取 ≥ 1 的整数推出结论时为True否则抛出jax.errors.InconclusiveDimensionOperationexport.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)) # jax._src.export.shape_poly.InconclusiveDimensionOperation: # Symbolic dimension comparison a 1 b is inconclusive.错误信息为文档示例。例如b 1、2 * a b 3可判定为True而b 2、a b、a - b 0无法判定。文档给出了四种应对策略若代码里用了内建max/min或np.max/np.min改用core.max_dim和core.min_dim把不等式比较推迟到编译时此时形状已知。用core.max_dim/core.min_dim改写条件分支例如把d if d 0 else 0写成core.max_dim(d, 0)。减少对维度必须是整数的依赖利用符号维度对多数算术运算 duck-type 为整数的特性例如把int(d) 5写成d 5。指定符号约束下一小节。一个细节值得注意当符号维度与非整数float、np.float、np.ndarray或 JAX 数组做算术时它会被隐式转成 JAX 数组结果可以当普通数组用但不能再用作 shape 中的维度。例如jnp.array(x.shape[0]) x可以导出但x.reshape(jnp.array(x.shape[0]) 2)会报TypeError: Shapes must be 1D sequences of concrete values of integer type。计算平均值这类常见写法jnp.sum(x, axis0) / x.shape[0]会自动走隐式转换导出与调用都正常。隐式约束与显式约束JAX 默认假设所有维度变量取值 ≥ 1并由此推导简单不等式如a 2 3、a * 2 1、a // 4 0。你可以通过改写 shape 规格加入隐式约束用2*b表示该维度是偶数且 ≥ 2用b 15表示该维度至少为 16。例如导出lambda x: x[0:16]时若规格是bJAX 需要验证切片大小不超过轴长而无法证明会失败改成b 15后导出成功_ export.export(jax.jit(lambda x: x[0:16]))( jax.ShapeDtypeStruct(export.symbolic_shape(b 15), dtypenp.int32))也可以用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))显式约束与隐式约束构成合取。文档说明当前对约束的推理能力有限收益最大的是变量与常数比大小的约束例如由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)。约束还能绕过推理规则的盲区。例如lax.slice_in_dim(x, 0, x.shape[0] % 3)在规格(b,)下会失败因为 JAX 无法证明b mod(b, 3)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)) # InconclusiveDimensionOperation: Symbolic dimension comparison # b mod(b, 3) is inconclusive.错误信息为文档示例。文档给出两条出路把规格改成3*b限制轴长为 3 的倍数JAX 即可把mod(3*b, 3)化简为0或者显式加上 JAX 正在试图证明的那条不等式b, export.symbolic_shape(b, constraints[b mod(b, 3)]) _ export.export(jax.jit(f))( jax.ShapeDtypeStruct((b,), dtypenp.int32))隐式和显式约束都在编译时通过同一机制检查见下文运行时形状断言错误。检查维度界与作用域export.symbolic_dim_bounds可以查看 JAX 能为某符号维度或由其推导的表达式证明的包含界。界是保守的、可能不紧无穷界表示 JAX 未能建立有限界并不证明该维度在数学上无上界batch, free export.symbolic_shape( batch, free, constraints(batch 128, batch 1024)) print(export.symbolic_dim_bounds(batch)) # (128, 1024) print(export.symbolic_dim_bounds(2 * batch 1)) # (257, 2049) print(export.symbolic_dim_bounds(free)) # (1, inf)以上输出为文档示例。约束存放在jax.export.SymbolicScope对象中每次调用export.symbolic_shape都会隐式创建一个新 scope。不要混合使用不同 scope 的符号表达式例如两次不同调用产生的a1与a2相加会报ValueError: Invalid mixing of symbolic scopes for linear combination。同一调用的产物共享 scope、可以自由互加要跨调用复用可以显式传 scopea, export.symbolic_shape(a,, constraints(a 8,)) b, export.symbolic_shape(b,, scopea.scope) # 复用 a 的 scope a b # 允许 # 也可以显式创建 scope my_scope export.SymbolicScope() c, export.symbolic_shape(c, scopemy_scope) d, export.symbolic_shape(d, scopemy_scope) c d # 允许另外注意 JAX trace 使用以 shape 为部分 key 的缓存打印结果相同的符号 shape 若属于不同 scope会被视为不同对象。维度变量必须能从输入形状解出当前唯一向已导出对象传递维度变量取值的途径是间接地通过数组参数的 shape。例如b的值在调用时从f32[b]类型第一个参数的 shape 推断这与 JIT 函数的调用约定一致。如果你的函数带一个决定程序内某些 shape 的整数参数如 top-k 的k直接把它符号化会失败因为k不出现在任何输入参数的 shape 中def my_top_k(k, x): # x: i32[4, 10], k 10 return lax.top_k(x, k)[0] x np.arange(40, dtypenp.int32).reshape((4, 10)) k, export.symbolic_shape(k, constraints[k 10]) export.export(jax.jit(my_top_k, static_argnums0))(k, x) # UnexpectedDimVar: Encountered dimension variable k that is not # appearing in the shapes of the function arguments错误信息为文档示例。若导出成功检查exp.in_avals/out_avals里k是否以符号形式出现调用时只需传非静态参数。文档给出的 workaround把参数k换成一个形状为(0, k)的空数组参数使k能从输入 shape 解出。第一维为 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]))文档示例 # 调用时必须构造并传入一个 (0, k) 形状的数组 exp.call(np.zeros((0, 3), dtypenp.int32), x) # Array([[ 9, 8, 7], ...], dtypeint32)文档示例另一类解不出来的情况维度变量出现在输入 shape 中但构成 JAX 目前解不了的非线性表达式例如规格(a * a,)会报ValueError: Cannot solve for values of dimension variables {a}只能解线性单变量约束。遇到这种错误时把规格改写成线性形式。运行时形状断言错误JAX 假设维度变量取严格正整数这一假设会在为具体输入形状编译时检查。例如符号输入 shape 为(b, b, 2*d)时JAX 会生成断言代码arg.shape[0] 1、arg.shape[1] arg.shape[0]、arg.shape[2] % 2 0、arg.shape[2] // 2 1。用不满足规格的具体形状调用时会在编译前的预处理阶段报错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)) # 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).错误信息为文档示例其中第三维 5 不满足2*d。这个错误信息本身可以当验证工具用它会逐条列出规格、已从各维度解出的维度变量取值指出哪条断言失败。调试 shape refinementshape refinement 在编译期为含维度变量或多平台的模块运行。出错时可用两个环境变量定位见 docs/export/export.md 的 Debugging 部分与 shape_poly 文档的 Debugging 部分# JAX_DUMP_IR_TO 指向一个目录导出模块会转储为 ..._export.mlir # shape refinement 出错时还会看到 refinement 前的 HLO 模块 # 文件名 ..._before_refine_polymorphic_shapes.mlir此时输入 shape 已静态化 # TF_CPP_VMODULErefine_polymorphic_shapes3 打开 shape refinement 各阶段的日志 JAX_DUMP_IR_TO/tmp/export.dumps/ TF_CPP_VMODULErefine_polymorphic_shapes3 python 你的触发导出或编译的脚本最后一行中你的触发导出或编译的脚本需替换为你实际运行导出/编译的 Python 脚本名。JAX_DUMP_IR_TO在文档其他位置也用于查看导出模块转储目录中还会同时出现 JIT 编译后的模块..._compile.mlir。可选序列化后在另一进程调用Exported对象本身可以在同进程内反复以不同具体形状调用。若要在另一进程或机器上编译执行、且不再需要 JAX 程序源码按 docs/export/export.md 的流程做两步先用export.export得到含 StableHLO 与调用元数据的Exported对象再用exp.serialize()序列化为 bytearrayflatbuffers 格式消费端用export.deserialize(serialized)还原后exp.call(...)serialized: bytearray exp.serialize() rehydrated_exp: export.Exported export.deserialize(serialized) res rehydrated_exp.call(np.ones((3, 4), dtypenp.int32)) # 任意匹配规格的 shape文档对此有明确警告序列化的 bytearray 必须是可信输入反序列化后执行它可能触发 jaxlib 中注册的任何 custom call。此外 export 有版本兼容窗口消费端 jaxlib 比导出端新不超过 6 个月、旧不超过 3 周跨版本部署时以两端构建时使用的 jaxlib 版本为准。限制与收尾核对完成一次符号 shape 导出后按下面的清单核对结果是否可用exp.in_avals/exp.out_avals中目标维度以符号表达式出现如int32[a, 2*b]说明规格生效用两个不同但都匹配规格的具体 shape 各调用一次确认都编译成功且结果与原生f(arg)一致——按文档的正确性保证导出成功时结果应当相同维度变量取值的唯一入口是输入 shape不要指望在调用时单独传参决定符号维度相等比较有 unsound 的边界b 1、a b也会返回False。因此if x.shape[0] ! 1: raise ...这类写法是 sound 的而if x.shape[0] ! 1: return 1这类依赖比较结果取值的写法不安全。更完整的规格语法floordiv、mod、max、min等维度表达式和函数签名见 jax/_src/export/shape_poly.py 中symbolic_shape与symbolic_args_specs的 docstring501 版本文档见 docs/501/shape-polymorphism.md。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表