资讯详情

资讯详情

JAX 原语到 MLIR 降级测试指南:深入解析 tests/filecheck 回归测试套件

人工智能机器学习深度学习编译器高性能计算【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址https://gitcode.com/GitHub_Trending/ja/jax点击查看免费下载导读JAX 的核心价值之一是把 Python NumPy 风格的算子lax原语降级lower为可供 XLA 编译器消费的 MLIR 中间表示进而编译到 CPU、GPU 与 TPU 等后端。tests/filecheck/目录维护着一套轻量级的 LLVM FileCheck 回归测试专门验证每个 JAX 原语降级后生成的 MLIR 文本是否符合预期。本文基于该目录下的 README 与全部测试源码系统讲解这套测试的设计动机、运行机制、CHECK 指令语法、各测试文件的覆盖范围以及其底层依赖的 MLIR 降级实现如jax._src.interpreters.mlir中的custom_call、merge_mlir_modules、lower_jaxpr_to_module等帮助你快速上手编写、运行与扩展这类测试并在改动 MLIR Python 绑定或 dialect 时第一时间捕获回归。这套测试的定位为什么需要 FileCheck 测试目录说明文档 给出了这套测试的明确定位目录中包含 LLVM FileCheck 测试用于验证JAX 原语可以被正确降级到 MLIR。其设计意图是提供一种快速且易于理解的回归捕获手段快速每个测试文件只聚焦少量算子直接对降级结果做文本匹配无需编译、无需 GPU/TPU 设备运行开销远小于完整 JAX 测试套件覆盖面精准专门捕捉两类典型回归——MLIR Python 绑定的变更如jax._src.lib.mlir中 IR 构造 API 的变化与JAX 所用各 MLIR dialect 的变更如 StableHLO、CHLO、arith、func 等 dialect 的语法或属性变化易读易维护预期结果以CHECK:注释直接写在生成 IR 的 Python 代码旁测试本身就是一份可读的降级结果快照。从测试的组织方式可以看出这是 JAX 在完整测试套件与直接审查源码之间特意保留的一层中间防线它不追求穷举所有算子行为那是tests/下lax_test.py、lax_numpy_test.py等数值测试的职责而是以最小的代价锁定降级产物长什么样。目录结构与运行机制tests/filecheck/目录共包含 9 个文件文件主题README.md目录定位与设计说明jax_filecheck_helpers.py提供print_ir辅助函数把任意 JAX 函数降级并打印 MLIRarray.filecheck.py数组结构与变形算子origami ops的降级math.filecheck.py逐元素数学算子的降级shapes.filecheck.py形状与 dtype 到 MLIR 类型系统的映射names.filecheck.py降级后 module / function 的命名规则custom_call.filecheck.pymlir.custom_call()各参数对 StableHLO 输出的影响subcomputations.filecheck.py子计算内联策略与多模块合并merge_mlir_modulesjax_mlir_ext.filecheck.pyJAX 自定义 C MLIR 扩展内联调用、traceback 定位、常量RUN 指令测试如何被驱动每个.filecheck.py文件头部都有一条RUN:注释例如# RUN: %PYTHON %s | FileCheck %s这条指令声明了测试的执行方式先用%PYTHONLLVM lit 测试框架注入的 Python 解释器运行当前脚本把脚本标准输出通过管道交给FileCheck工具FileCheck再根据脚本源码中的CHECK:注释逐行验证输出。也就是说测试脚本本身就是期望值生成器——脚本先打印出真实的降级 IR同时内嵌的CHECK注释定义了期望的模式两者在同一份文件里形成自洽的校验对。jax_mlir_ext.filecheck.py稍有不同其 RUN 行带上了-dump-inputalwaystests/filecheck/jax_mlir_ext.filecheck.py#L15用于在匹配失败时打印完整输入行便于调试复杂的多行模式。CHECK 指令语法速查FileCheck的指令都是CHECK前缀注释本套件实际用到的主要有CHECK:—— 按顺序出现的普通匹配行CHECK-LABEL:—— 标记一个测试段的起始用于把不同测试块的匹配范围隔离开每个print_ir用例前都有一个TEST: xxx输出行与之对应CHECK-SAME:—— 要求匹配内容必须与上一行同一物理行的后续部分常用于验证 IR 类型签名续行CHECK-NOT:—— 断言接下来的内容中不出现指定模式如subcomputations.filecheck.py用CHECK-NOT: func private cumsum验证子计算只生成一次CHECK-DAG:—— 允许匹配顺序不固定的行如math.filecheck.py中lax.conj的hlo.real/hlo.imag/hlo.neg/hlo.complex四行降级顺序不保证{{...}}—— 正则表达式转义如{{[0-9]}}匹配任意位数用于行号这类不稳定内容。print_ir把 JAX 函数变成 FileCheck 的输入所有基于print_ir的测试都依赖 jax_filecheck_helpers.py 中的辅助函数其实现只有约 15 行却串起了原型数组 → 真实输入 → jit 降级 → 打印 IR的完整链路def print_ir(*prototypes): def lower(f): inputs jax.tree.map(np.array, prototypes) flat_inputs, _ jax.tree.flatten(inputs) shape_strs .join([f{x.dtype.name}[{,.join(map(str, x.shape))}] for x in flat_inputs]) name f.func.__name__ if hasattr(f, func) else f.__name__ print(f\nTEST: {name} {shape_strs}) print(jax.jit(f).lower(*inputs).compiler_ir()) return lower关键点解读原型数组prototypes调用方传入的np.empty([2, 7], np.int32)这类形状 dtype 模板并不直接参与计算而是被jax.tree.map(np.array, ...)原样转成真实输入数组只用来决定参数的类型与形状TEST 标签行先打印\nTEST: 函数名 dtype1[shape1] dtype2[shape2]与各用例的CHECK-LABEL: TEST: ...精确对应既起到分段作用也让人一眼看出这个用例在测什么形状什么类型降级入口jax.jit(f).lower(*inputs).compiler_ir()完成追踪 → jaxpr → MLIR 模块的全流程。compiler_ir()定义于 stages.py默认返回 StableHLO dialect 的 MLIR 模块对象print即输出其文本形式。print_ir支持两种用法直接调用print_ir(prototype)(partial(lax.xxx, ...))或作为装饰器包裹自定义函数如names.filecheck.py中print_ir(...) jax.jit def foo(x): return x 2。装饰器写法还能配合partial、闭包等构造更复杂的被测函数。array.filecheck.py数组结构算子的降级array.filecheck.py 文件头注释点名了主题Tests for lowering of array origami ops into MLIR.数组折纸算子覆盖 12 个典型的形状/结构变换算子每个用例都精确断言了目标算子名与结果张量类型被测算子期望的 MLIR 算子期望类型签名lax.concatenatedimension1hlo.concatenatetensor2x12xi1lax.broadcast_in_dimhlo.broadcast_in_dimtensor3x2x5x7x2xi1lax.iotahlo.iotatensor10xf32lax.padhlo.padtensor11x52xi32lax.reduce_sumaxes(0,2)hlo.reducehlo.addtensor3xi32lax.reshapehlo.reshapetensor42xi32lax.revhlo.revtensor2x7xi32lax.selecthlo.selecttensor2x7xi1, tensor2x7xi32lax.sorthlo.sorttensor2x7xi32lax.squeezehlo.reshapetensor2x7xi32lax.top_kchlo.top_ktensor2x7xi32lax.transposehlo.transposetensor7x2xi32几个值得注意的细节同一原语可降级为不同算子lax.squeeze并没有专属的hlo.squeeze而是复用hlo.reshape这正体现了测试的价值——锁定的不是JAX 算子的语义而是实际发往编译器的算子形态dialect 并非只有 StableHLOlax.top_k降级为 CHLO dialect 的chlo.top_k。CHLO 是 StableHLO 之上的高阶算子层会在后续编译阶段被继续展开测试因此也间接验证了 JAX 与 CHLO dialect 绑定的兼容性pad的参数编码padding_config((2, 3, 4), (4, 5, 6))表示每个维度上的 (前填充, 后填充, 内部间隔)结果223行、745列类型断言tensor11x52xi32与之完全一致说明测试同时对参数语义做了校验。math.filecheck.py逐元素算子的降级全景math.filecheck.py 是本套件中用例最密集的文件覆盖约 60 个逐元素算子并特意在文件开头启用了 x64 模式jax.config.update(jax_enable_x64, True)以便覆盖f64类型。从这些用例可以归纳出 JAX 降级映射的几条规律基础算术/位运算直接映射 StableHLOlax.add → hlo.add、lax.mul → hlo.mul、lax.div → hlo.div、lax.rem → hlo.rem、lax.pow → hlo.power、lax.neg → hlo.negate、lax.bitwise_and → hlo.and、lax.bitwise_or → hlo.or、lax.bitwise_xor → hlo.xor、lax.bitwise_not → hlo.not、lax.shift_left → hlo.shift_left、lax.shift_right_arithmetic → hlo.shift_right_arithmetic、lax.shift_right_logical → hlo.shift_right_logical基础超越函数也走 StableHLOsin/cos/tan/sqrt/exp/log等映射为hlo.sin、hlo.cos、hlo.tan、hlo.sqrt、hlo.exp、hlo.log等较复杂或较新的数学函数映射到 CHLOlax.acos → chlo.acos、lax.acosh → chlo.acosh、lax.asinh → chlo.asinh、lax.atan → chlo.atan、lax.atanh → chlo.atanh、lax.bessel_i1e → chlo.bessel_i1e、lax.cosh → chlo.cosh、lax.sinh → chlo.sinh、lax.digamma → chlo.digamma、lax.erf/erfc/erf_inv → chlo.erf/erfc/erf_inv、lax.lgamma → chlo.lgamma、lax.nextafter → chlo.next_after有些数学上等价的映射出乎意料lax.asin降级为hlo.atan2用atan2(x, sqrt(1-x²))的公式实现lax.integer_pow展开为多次hlo.mul用CHECK-DAG匹配lax.conj则由hlo.realhlo.imaghlo.neghlo.complex组合而成比较运算携带比较语义lax.eq/ge/gt/le/lt/ne统一映射为hlo.compare EQ/GE/GT/LE/LT/NE且比较方向标签FLOAT / SIGNED / UNSIGNED由操作数类型决定——例如f32比较断言FLOAT、i64断言SIGNED、ui16断言UNSIGNED同时complex类型也标记为FLOAT见eq complex128[]用例类型转换的多样性lax.convert_element_type依据源/目标类型走完全不同的降级——浮点转浮点用hlo.convert、复数转实浮点用hlo.real、浮点转布尔用hlo.compare与零比较。这些用例同时验证了 JAX 的标量 dtype 表示int32 → i32、uint32 → ui32、bool → i1、float32 → f32、bfloat16 → bf16、complex64 → complexf32构成了 shapes.filecheck.py 的类型映射专题的基础。shapes.filecheck.py 与 names.filecheck.py类型与符号命名契约JAX 类型到 MLIR 的完整映射shapes.filecheck.py 专门验证JAX shapes and types降级到 MLIR 类型系统的正确性覆盖了布尔与零维标量bool[7] → tensor7xi1有符号整数全系列i8 / i16 / i32 / i64含零长度维度tensor0xi16与多维度tensor2x3x4xi64无符号整数全系列ui8 / ui16 / ui32 / ui64含含零维度tensor4x0x1xui8浮点全系列f16 / bf16 / f32 / f64复数complex64 → tensorcomplexf32complex128 → tensorcomplexf64。其中cos complex64[]与cos complex128[]两个用例带 TODO 注释when the accuracy of lax.cos is fixed upstream, undo relevant parts of jax PR 19823其期望结果是hlo.cosine直接作用于复数输入从源码结构看这提示复数cos的精度问题曾导致 JAX 侧对上游做了针对性 workaround而此类上游修复后需回退的标记也正是 FileCheck 测试承担变更追踪任务的典型场景。dtype 到 IR 类型的底层转换实现在 mlir.py 的dtype_to_ir_type它通过_dtype_to_ir_type查表分发未知 dtype 会抛出TypeError——这也解释了为什么测试里出现的每个 dtype 都必须有稳定的映射。模块与函数命名契约names.filecheck.py 只含两个用例却锁定了 JAX 降级产物在符号层的命名契约# CHECK-LABEL: TEST: neg int32[7] # CHECK: module jit_neg # CHECK: func public main以及装饰器 jax.jit组合下module jit_foo的命名。这意味着每个jax.jit降级结果都是一个以jit_函数名命名的 MLIR 模块其公开入口统一命名为func public main。这两个用例的存在非常关键——JAX 的多模块合并、跨模块调用、序列化等机制都依赖这套稳定的符号命名任何破坏该契约的改动例如绑定层把main改名、把公开函数改成私有都会被立即捕获。custom_call.filecheck.pycustom_call 的稳定输出契约custom_call.filecheck.py 直接构造 MLIR IR 并调用 mlir.custom_call 这一底层接口验证其各参数如何反映到最终 StableHLO 文本中。测试定义了一个print_custom_call辅助函数先通过mlir.make_ir_context()创建 IR 上下文用ir.RankedTensorType.get(shape, mlir.dtype_to_ir_type(dtype))把ShapedArray转成 MLIR 类型再构建func函数体并在其中调用mlir.custom_call最后module.operation.verify()校验合法性并打印。六个用例精确锁定了参数到属性的映射参数组合生成的属性默认参数api_version 2 : i32backend_config api_version1, has_side_effectTruehas_side_effect true且api_version降为 1backend_configbhellobackend_config hellocalled_computations[a, b]called_computations [a, b]operand_output_aliases{1: 0}output_operand_aliases [#stablehlo.output_operand_alias...]operand_layouts/result_layoutsoperand_layouts [dense[0, 1]...]result_layouts [...]对照 mlir.py 的实现可以看到更多细节backend_config支持str/bytes/dict三种形式其中dict形式会触发一个特殊路径——由于 StableHLO 的CustomCallOp构造函数要求backend_config必须是字符串属性dict内容会被存到未注册的mhlo.backend_config属性中且必须使用api_version1才能被正确处理operand_output_aliases则会被编码为OutputOperandAlias属性且当只有一个输出时输出元组索引自动置空见 mlir.py。这类构造函数限制导致的实现变通正是 FileCheck 测试需要长期锁定的不稳定区域。subcomputations.filecheck.py内联策略与模块合并subcomputations.filecheck.py 验证两类 MLIR 辅助能力是理解 JAX 编译器后端如何组织代码的关键子计算的内联控制第一个用例cumsum_only_once体现了 JAX 对每个形状只生成一次子计算的保证lax.cumsum的降级被标注为inlineFalse因此同一函数里对两个同形状数组分别调用cumsum只会生成一个func private cumsum# CHECK: func private cumsum # CHECK-NOT: func private cumsumCHECK-NOT在这里的作用是断言第二次出现即失败从而保证重复使用同一形状的子计算不会产生重复的私有函数。这与 JAX 的常量折叠/子计算复用机制直接相关是防止代码膨胀的重要契约。多模块合并与符号重命名后两个用例测试mlir.merge_mlir_modules。第一个用例先把两个jax.jit函数的降级模块make_module(10)与make_module(20)取出将后者的文本在前者上下文中重新ir.Module.parse后合并期望结果中出现func public main、func private f、func private m2_main_renamed、func private f_0——注意f与main都被自动重命名以避免冲突第二个用例则直接解析两段完全相同的文本模块并合并期望得到f、f_0、f_1三个函数其中f_1同时被main调用。这些期望行为与 merge_mlir_modules 的实现逐行对应该函数要求dst_module与src_module共享同一 IR 上下文assert dst_module.context src_module.context遍历源模块中所有func符号凡是与目标符号表冲突或名为main的一律重命名为base_ii从 0 递增确保新名字既不在源模块也不在目标模块中随后统一把源函数设为private可见性用replace_all_symbol_uses重写所有引用最后把源模块的操作整体搬进目标模块并返回main重命名后的真实名字。测试锁定的正是这套alpha 重命名 私有化 符号表重写的完整流程。jax_mlir_ext.filecheck.pyC 扩展的 MLIR 绑定jax_mlir_ext.filecheck.py 是唯一直接测试 JAX 自带 C 扩展jaxlib.mlir._mlir_libs._jax_mlir_ext的文件覆盖三个能力inlined_func_call函数内联调用通过callee私有、含两个算术 op与caller公开两个函数构造模块在 caller 中调用扩展的inlined_func_call把 callee 的运算内联进 caller再运行builtin.module(symbol-dce)pass 清理死符号最后带 debug info 打印。期望输出展示了内联后 caller 主体只含stablehlo.add与stablehlo.multiply两个 op以及一整套由loc(...)构成的调用栈位置链callee_stack、caller_stack、callsite(...)嵌套。这是 JAX 实现函数内联优化时所依赖的原生绑定测试确保绑定层的位置信息与符号处理不被破坏TracebackToLocationCachetraceback → MLIR location抓取真实的 Python traceback通过code_to_filename回调把代码对象映射为文件名再以frame_limit1000与frame_limit2两个档位打印loc(callsite(...))嵌套结构。正则{{[0-9]}}用来容忍行号变化验证的是调用栈到 MLIR location 的映射层级这一绑定能力arith_constantarith 常量验证布尔arith.constant true、整数arith.constant 42 : i32、浮点arith.constant 3.140000e00 : f32的构造与打印直接关系到 JAX 在 MLIR 中生成常量时的底层正确性。这些用例还展示了文件头jax.config.parse_flags_with_absl()的用法——该文件是套件中唯一解析 absl 标志的测试说明它具备独立的运行配置需求。如何运行与扩展这套测试运行方式FileCheck是 LLVM 项目自带的测试工具JAX 仓库中的 CI 脚本ci/run_pytest_cpu.sh 等在带 MLIR 绑定的构建中会将其纳入测试流程。本地手动运行某个用例的等价命令是python tests/filecheck/math.filecheck.py | FileCheck tests/filecheck/math.filecheck.py即运行脚本 → 管道输出 → 用源码中的 CHECK 注释校验。在 LLVM lit 集成环境中则由 RUN 行中的%PYTHON与FileCheck两个占位符自动完成同一件事。新增测试用例的通用范式参照print_ir用法为一个新算子或新重载添加降级测试只需三步选原型用np.empty([...], np.dtype)/jnp.bfloat16(0)等构造形状 类型的模板输入写 CHECK在调用前写好CHECK-LABEL: TEST: 名字与print_ir自动打印的 TEST 行一致以及CHECK: 期望算子、CHECK-SAME: tensor...断言类型签名跑通并核对先直接运行脚本查看实际 IR再把关键算子名 结果类型 关键属性固化为 CHECK 行——注意只锁定稳定契约对行号、SSA 变量名等易变内容用{{...}}正则或CHECK-DAG容忍。如果被测逻辑需要直接操作 MLIR IR而非从 JAX 函数出发则参照custom_call.filecheck.py与jax_mlir_ext.filecheck.py的模式用mlir.make_ir_context()ir.Module.createInsertionPoint手工构建模块最后module.operation.verify()并打印。小结把降级正确性变成可回归的契约JAX 的编译管线可以概括为Python 函数 → jaxpr → MLIRStableHLO/CHLO→ 后端可执行程序其中 lower_jaxpr_to_module 是连接 jaxpr 与 MLIR 的关键枢纽。tests/filecheck/这套测试的价值在于把这条管线的中间产物文本变成了可机器校验的契约算子名、dialect 归属hlo.* vs chlo.*、类型签名tensor2x12xi1这类、参数编码pad 配置、比较方向标签、符号命名module jit_*/func public main、内联与合并策略、custom_call 属性序列化以及 C 扩展的 IR 构造行为都被逐条固化下来。无论是 MLIR Python 绑定的重构、dialect 语法的升级还是 JAX 侧降级逻辑的调整只要运行这批轻量测试就能在数秒内定位哪些算子的降级产物发生了非预期变化从而在完整测试套件暴露数值问题之前先行守住编译链路的底层正确性。赞分享人工智能机器学习深度学习编译器高性能计算【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址https://gitcode.com/GitHub_Trending/ja/jax点击查看免费下载相关推荐用 LLVM FileCheck 守护 JAX 原语到 MLIR 的降级JAX 轻量级降级回归测试指南用 LLVM FileCheck 守护 JAX 原语到 MLIR 的降级JAX 轻量级降级回归测试指南 JAX 的核心工作方式之一是将 PythonNum机器学习深度学习radare2 代码风格回归测试指南深入解析 sys/clang-format-radare2 与 test/indent 测试套件radare2 代码风格回归测试指南深入解析 sys/clang format radare2 与 test/indent 测试套件 本指南以 test/in逆向工程网络安全Haxe 杂项测试套件解析tests/misc 编译回归测试的约定、实现与本地运行指南Haxe 杂项测试套件解析tests/misc 编译回归测试的约定、实现与本地运行指南 Haxe 编译器的开发离不开一套覆盖面极广的回归测试体系其中位于仓库编程语言编译器语言运行时标准库上一篇Foundry Anvil 性能优化解析eth_getBlockReceipts 大区块批量收据加速下一篇sonic 的 JSON 兼容性基准深入解读 JSONTestSuite 与 RFC 8259 边界用例创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
觉得有用,分享给同行:

为您的企业打造数字门面

稳重轻奢商务风格,端正雅致视觉,长效耐看不易过时。

立即咨询 →