Equivalence Checking of ML GPU Kernels

发表时间: 2025-11 · arXiv:2511.12638 (Stanford, Microsoft Research India, Google DeepMind)

原文: https://arxiv.org/abs/2511.12638

作者/机构: BENJAMIN DRISCOLL (Stanford University), KSHITIJ DUBEY (Microsoft Research), ANJIANG WEI (Stanford University), NEERAJ KAYAL (Microsoft Research), RAHUL SHARMA (Google DeepMind), ALEX AIKEN (Stanford University)

主要贡献

随着深度学习和大型语言模型(LLMs)的快速发展,企业在执行GPU计算内核上花费了巨大的成本。因此,这些内核成为了激进优化的主要目标。近年来,越来越多的工作利用LLMs来生成GPU内核(例如NVIDIA使用DeepSeek-R1,Stanford研究人员使用OpenAI o3等),但这些方法对生成的内核没有提供任何形式化的正确性保证。给定两个GPU内核——一个参考实现和一个优化后的对应实现,本文的核心研究目标是证明它们的扩展等价性(extensional equivalence),即证明它们在所有有效输入上产生相同的输出,从而验证优化实现的正确性。

由于GPU编程引入了海量线程之间细粒度同步的独特挑战,且优化通常会重排浮点算术操作导致微小的数值差异,目前文献中尚未存在针对GPU程序的等价性检查器。本文提出了首个针对GPU内核的等价性检查器,并用它形式化验证了由人工、LLM和编译器优化的机器学习(ML)内核的正确性。

本文的创新点与主要贡献包括:
1. 提出了一种线性时间算法,用于检测使用屏障同步(barrier synchronization)的特定类别GPU内核(结构化CTA,structured-CTAs)中的数据竞争(data races)和死锁(deadlocks)。
2. 证明了形式为 $\sum_{i=1}^{m} p_{i}(\bar{x}) e^{h_{i}(\bar{x})} = 0$ 的等式的可判定性(decidability),其中 $p_{1}, h_{1}, \ldots \in \mathbb{R}[\bar{x}]$ 为多元多项式。
3. 开发了首个针对实用GPU内核类别的等价性检查器实现——Volta。Volta能够验证优化后的卷积、矩阵乘法、归约和注意力机制的正确性。作者证明了Volta保证了数据竞争自由,并满足真阳性属性(true positives property),同时在Agda证明助手中形式化了关键的组成引理。

背景知识与设计原则

动机示例分析:为了说明问题,考虑大型语言模型(LLMs)推理中的主要性能瓶颈——注意力机制中的softmax计算。朴素的参考实现(如CUDA代码)通常需要使用共享内存(shared memory)来存储中间指数值,并依赖 syncthreads 屏障原语来同步所有线程,确保所有共享内存的写入在任何线程继续之前都是可见的。而FlashAttention 【21, FlashAttention: Fast and MemoryEfficient Exact Attention with IO-Awareness, 2022 NeurIPS】中使用的在线softmax计算则避免了使用共享内存,通过维护运行时的最大值和归一化因子来实现数值稳定和内存高效的流式处理。这两种实现由于浮点舍入差异在位级别(bit-wise)上并不完全相同,但都旨在计算相同的数学表达式。概念上,等价性分析将softmax内核展开为每线程程序的列表(如包含4个线程的CTA),每个线程执行不包含分支或循环的直线型代码,且所有内存访问的地址都是静态已知的。通过执行简单的轮询调度(round-robin scheduling),分析器记录每个线程读取或写入的共享内存位置,若缺少屏障同步则报告数据竞争。如果没有数据竞争或死锁,符号执行器会为朴素实现和优化实现分别生成输出张量元素的符号表达式(包含实数上的指数和多项式运算),最后通过决策过程验证这些符号等式是否成立,从而得出两个程序是否语义等价的结论。
图1. 朴素softmax(左)需要在共享内存中具体化所有指数运算;在线softmax(右,FlashAttention中使用)维护一个运行时的最大值和归一化因子,从而实现数值稳定且内存高效的流式处理。

GPU同步操作机制:NVIDIA GPU支持独立线程调度,在缺乏显式同步的情况下,每个线程独立执行。基础同步原语是 syncthreads,它是一个跨越协作线程阵列(CTA,即线程块)内所有线程的屏障。另一个同步原语是 syncwarp,它在单个warp(32个线程)内的用户指定子集上运行,常用于张量核心(tensor cores)编程(如 wmma::mma_sync)。大多数先前的GPU内核验证技术仅支持 syncthreads,无法推理依赖warp级别同步的ML内核。本文考虑的同步机制是阻塞可用线程子集的屏障,不包括异步的生产者-消费者风格原语。

为什么选择结构化CTA(Structured-CTAs):为了使等价性检查可行,本文假设被分析的程序是结构化CTA。程序约束条件:要检查等价性,程序必须在已知数量的CTA和每个CTA的线程数上运行;每个CTA必须在大小静态已知的张量块上运行;在每个CTA的每个线程中,给定线程ID(tid),所有分支目标和内存访问中使用的地址必须是静态已知的;由于是CTA对CTA进行比较,不支持混合CTA职责的优化等价性检查。假设的合理性:这些假设反映了实现高GPU性能的既定最佳实践。例如,JAX框架仅允许固定维度的张量;高效的ML内核通常不做出真正的数据依赖的运行时决策。在这些假设下,每个单独线程执行的代码可以表示为直线型(无动态分支或循环)。不适用的内核:诸如 top-$k$/argsort 等控制流依赖于数据的内核不属于结构化CTA类别。浮点数处理原则:由于IEEE-754标准甚至不满足结合律,严格基于浮点算术的等价性检查会拒绝许多激进优化的内核。因此,Volta遵循与现有优化器相同的约定,将张量元素建模为实数(reals)进行处理。
图3. 结构化CTA语言$\mathcal{L}$的语法。

方法细节

形式化设置

结构化CTA语言$\mathcal{L}$的定义:为了形式化分析,作者开发了一种简单的结构化CTA语言 $\mathcal{L}$,它捕获了PTX(NVIDIA GPU的底层汇编)的相关方面。$\mathcal{L}$ 程序的语法由寄存器 $r$、常量 $c$、共享内存地址 $g$ 以及线程ID集合 $I$ 组成。作者通过一对小步语义关系为 $\mathcal{L}$ 赋予了动态特性:顶层配置 $( \mathcal{G}, \mathcal{R}, P )$ 包含共享内存映射 $\mathcal{G}$、寄存器文件映射 $\mathcal{R}$ 以及当前程序 $P$。
调度与同步规则Schd 规则允许任何未阻塞的线程推进,模拟机器的非确定性执行。Sync 规则推进阻塞在同步屏障 $I$ 处的线程,前提是 $I$ 中的所有线程都处于 sync Ireturn 状态。当 $I = {1, \ldots, N}$ 时,它模拟 syncthreads;当 $I$ 是warp子集时,它模拟 syncwarpmma.sync。如果两个线程在不同的集合上同步且存在交集,程序将陷入死锁(stuck)。
线程级执行规则:线程级小步关系定义了 ConstBinOpRdRegRdMemWrMem,用于更新内存 $\mathcal{G}$ 和寄存器 $R$。算术运算符 $\oplus$ 被建模为黑盒纯函数,隐藏了内部可能的控制流。

符号执行

数据竞争定义与上下文扩展:当两个内存访问中至少有一个是写入,且存在两种执行轨迹使得这两个访问的顺序可以互换时,就会发生数据竞争。为了在符号执行中检测竞争,作者将 $\mathcal{L}$ 的动态特性扩展为在发生竞争时返回错误值 $\perp$。具体而言,引入了一个额外的上下文 $\chi$,用于跟踪每个地址 $g$ 的内存事件(MemEvs):包括每个读取该地址的线程及其尚未同步的线程集合,以及最后写入该地址的线程及其尚未同步的线程集合。
竞争检测谓词:定义了谓词 noRacingRd(检查当前写入线程是否与所有先前的读取线程同步)和 noRacingWr(检查当前读取/写入线程是否与最后的写入线程同步)。
规则的细化:线程级读取规则 RdMem' 和写入规则 WrMem' 被修改为在执行内存操作前检查上述谓词,并相应更新上下文 $\chi$。Sync' 规则被修改为在同步发生时清除 $\chi$ 中对应的未同步记录。如果谓词检查失败,则通过 RdMemBadWrMemBad 规则直接转移到错误状态 $\perp$。

形式化保证

合流性(Confluence)证明:等价性检查的健全性依赖于处理符号执行中调度引入的非确定性。作者在Agda证明助手中验证了核心引理(Lemma 1):对于所有程序,所有可能的调度要么产生相同的输出状态,要么报告数据竞争($\perp$),要么陷入死锁。这是通过证明自反闭包满足单步菱形属性(diamond property)来实现的,其中大部分情况归结为不同同步操作的交换律。
数据竞争与死锁检测的理论保证
- 证明了未检查的动态轨迹与带有竞争检查的动态轨迹之间的对应关系(Lemma 2和Lemma 3)。
- 证明了初始内存和寄存器值不影响竞争检测(Lemma 4)。
- 数据竞争自由(定理1):如果程序存在数据竞争,带有检查的动态必定会评估为 $\perp$。
- 真阳性属性(定理2):如果带有检查的动态评估为 $\perp$,则程序确实存在数据竞争。这是通过将导致 $\perp$ 的轨迹重新排序,使得冲突的内存访问相邻,从而构造出两个顺序相反的真实轨迹来证明的。
- 死锁检测(定理3与定理4):由于合流性,如果任何轨迹从某个状态死锁,则所有轨迹都会死锁;如果未检查的动态死锁,符号执行必定会死锁或报告竞争。
等价性检查的健全性与完备性
- 健全性(定理5):如果决策程序判定两个程序在符号输入上产生相等的项,那么在未检查的动态下,对于所有实数输入,这两个程序都会评估出相同的结果。这依赖于符号执行作为“健全抽象”(Lemma 5)。
- 完备性(定理6):如果决策程序判定符号输出不相等,则必然存在某个实数映射,使得两个程序在未检查的动态下产生不同的输出。

决策程序

支持的数学运算:决策程序支持包含实数加法、乘法和指数运算的表达式等式判定。虽然未形式化给出 max 和 min 的说明,但实现中通过规范化处理了这些操作(例如在softmax的最大值减法中)。
指数多项式等式的可判定性:作者证明了形式为 $f_1(x_1, \ldots, x_n) = f_2(x_1, \ldots, x_n)$ 的等式是可判定的,这等价于证明 $f_1 - f_2 \equiv 0$。
- 单变量定理(定理7):假设 $\mathbb{R}[x]$ 上的多项式 $p_i(x)$ 和 $h_i(x)$ 满足 $h_i(x)$ 互不相同,若 $\sum_{i=1}^{m} p_{i}(x) e^{h_{i}(x)} = 0$,则必定所有 $p_i(x) = 0$。证明过程通过假设最大首项系数的指数,并在 $x \to \infty$ 时取极限,得出对应多项式必为零的矛盾,从而通过归纳法完成证明。
- 多变量推论(推论8):利用Schwartz Zippel引理,通过随机代入 $x_i = \alpha_i y$,将多变量多项式以高概率转化为单变量多项式,从而证明了如果 $\sum_{i=1}^{m} p_{i}(\bar{x}) e^{h_{i}(\bar{x})} = 0$,则所有 $p_i(\bar{x}) = 0$。这避免了现有文献中依赖代数数的复杂数论机制,直接处理任意实数系数。

系统实现 (Implementation)

Volta架构与工作流:Volta是用Rust编写的概念验证工具。它接受PTX ISA 9.1编写的CTA代码,以及网格维度、张量位置和大小等参数。它由两个主要组件构成:执行符号执行生成验证条件(VCs)的模拟器,以及解析这些条件的决策程序。
符号执行的实现细节
- 在符号执行前,PTX代码被降级以解析变量、参数、寄存器名和标签。
- 模拟器为每个被写入的共享或全局内存单元跟踪一个表达式和上下文 $\chi$(使用固定宽度位集实现)。表达式作为4字节索引传递给一个bump分配器,该分配器保存16字节的表达式表示。
- 常量折叠与化简:在构建表达式时急切地执行常量折叠,包括化简 $\min(\infty, x)$、$\max(-\infty, x)$ 等,浮点常量使用GMP库的任意精度有理数表示。这使得所有分支和内存访问在发生时能被静态解析。
- 调度策略:线程以轮询方式被调度,直到阻塞在屏障或退出。如果所有线程都阻塞,调度器检查是否可以释放屏障;如果不能,则报告死锁。
决策程序的实现细节
- 规范化器(canonicalizer)将表达式转换为多项式的比例(有理函数)。多项式由排序的项和系数列表组成,项由符号输入、未解释函数或有理数的min/max组成。
- 所有内容都被哈希共用(hash-consed)以提供稳定的排序顺序。
- 缓存与引用计数:规范化的结果基于表达式的生成时间索引被记忆化/缓存。为了减少内存占用,只有在表达式被多个其他表达式引用时才进行缓存,引用计数通过遍历bump分配器确定。
- 支持的PTX操作:支持 bar.syncbar.warp.syncshfl.syncactivemaskldmatrixmma.syncwmma。符号化支持加、乘、除、倒数、融合乘加、指数、最大值/最小值、浮点宽度转换、饱和和ReLU截断。不支持平方根或三角函数。

实验环境

  • 硬件配置:测试平台为配备128核 2.25 GHz AMD EPYC 7742 处理器和 995 GB RAM 的服务器。用于验证未记录的硬件越界行为的GPU型号包括NVIDIA Ampere、Hopper和Blackwell架构。
  • 软件配置:Volta工具由Rust语言实现,依赖GMP库处理任意精度有理数。基线对比工具使用Z3 SMT求解器(版本4.8.12)。分析的汇编代码版本为NVIDIA PTX ISA 9.1。
  • 数据集/内核任务:验证对象涵盖四大类:
    1. 人工生成的内核:包括7个版本的归约(Reduction)内核【30, Optimizing Parallel Reduction in CUDA】、7个版本的矩阵乘法(MatMul)内核【11, How to Optimize a CUDA Matmul Kernel】以及多种注意力机制(Attention、FlashAttention-1/2、因果掩码注意力)。
    2. LLM生成的内核:由Stanford研究人员生成的2D卷积内核,以及由Claude Code智能体生成的矩阵乘法内核。
    3. 编译器生成的内核:由TileLang编译器生成的不同分块大小的矩阵乘法内核(调用CUTLASS模板)。
    4. 包含已知Bug的GitHub开源内核:来自OpenMM和Megatron-LM的真实数据竞争修复提交。

实验结果

  1. 人工生成的内核验证
    - 归约内核(Reduction):表1展示了验证结果。Volta成功识别出Red-5、Red-6和Red-7中由于使用了被弃用的warp同步而导致的数据竞争,并在符号执行阶段拒绝了它们(检测时间 $\le 0.0011$ 秒)。同时,Volta证明了Red-1到Red-4与基线Red-1的等价性。
    - 矩阵乘法内核(MatMul):表2展示了验证结果。所有7个优化版本(包含合并内存访问、共享内存缓存、消除bank冲突、多级tiling等优化)均顺利通过数据竞争检查,并被成功证明与基线MatMul-1等价,VC求解时间在89到94秒之间。
    - 注意力机制(FlashAttention):表3和表4展示了结果。Volta成功证明了采用在线softmax和张量核心的FA1、FA1-TC和FA2-TC与朴素Attention实现的等价性。由于需要协调朴素softmax与在线softmax累加的两种不同有理函数表示,非掩码版本的VC求解时间较高(约154秒)。而因果掩码版本(Causal-Attention)由于过滤方式的差异在符号执行中被消除,比较的表达式完全相同,VC求解时间极短($\le 0.13$ 秒)。

  2. LLM生成的内核验证
    - 2D卷积(Conv2D):表5展示了验证结果。Volta被用于验证经过13轮LLM优化的2D卷积内核(转换为隐式GEMM并使用张量核心)。在验证过程中,Volta发现了一个未记录的硬件行为:代码中存在对共享内存的越界读取,尽管在Ampere、Hopper和Blackwell硬件上不会崩溃并静默返回0(甚至NVIDIA Compute Sanitizer也未报错),但Volta正确捕获了该越界行为并抛出异常。修复该问题后,内核被证明等价。
    - Claude Code生成的MatMul:表6展示了结果。Claude Code自动生成的三个内核(GEMM-1/2优化了共享内存,GEMM-3使用了张量核心)均被Volta成功证明与基线参考实现等价。

  3. 编译器生成的内核验证
    - TileLang矩阵乘法:表7展示了结果。验证了由TileLang编译器生成的不同输出块大小(32x32x32至64x64x32)的优化内核(调用CUTLASS模板并使用张量核心)与参考实现的等价性。结果表明,Volta的运行时间随着分块大小的增加而优雅扩展(最大分块VC生成18.6秒,求解60.4秒)。

  4. 数据竞争检测基线对比
    - 对比FaialAA:表8展示了与最先进的GPU数据竞争检查器FaialAA【46, Sound and Partially-Complete Static Analysis of Data-Races in GPU Programs, 2024 OOPSLA2】的对比。FaialAA由于无法跟踪具体值,会将类似 M[tid]=tid; x=M[tid]; M[x]=... 的代码误报为数据竞争。Volta通过符号执行积极传播常量,避免了此类假阳性。此外,Volta成功检测出了OpenMM和Megatron-LM历史提交中真实存在的数据竞争,并证明了修复后代码的竞争自由。

  5. 等价性检查基线对比
    - 对比Z3求解器:表9展示了单元素VC求解时间的对比。对于不包含指数的18个基准测试(归约、MatMul、Conv2D),Z3能够成功求解。但对于包含softmax指数运算的注意力机制基准测试,即使添加了指数乘积公理($e^x e^y = e^{x+y}$),Z3也会超时(>300秒)或返回Unknown。相反,Volta得益于缓存机制和专门的规范化过程,其中位数单元素求解时间极低(例如Attention仅需1.3毫秒,FA系列约90毫秒),远胜于现成的SMT求解器。

补充细节 (Related Work)

  • 等价性检查相关工作:现有的等价性检查工作(如Alive针对位向量,Mirage针对多线性算子)缺乏对GPU并行性或同步的支持。Mirage 【76, Mirage: a multi-level superoptimizer for tensor programs, 2025 OSDI】使用李和吴【43, Identity Testing for Circuits with Exponentiation Gates, 2026 ITCS】的恒等测试程序处理包含指数的条件,但该算法缺乏本文所提供的可证明的健全性。此外,完全模拟IEEE-754语义会导致张量计算重排优化的等价性检查失败。
  • 数据竞争检测相关工作:传统数据竞争分析只考虑内存访问发生的位置而不考虑具体数据。FaialAA是健全数据竞争检查的最先进技术,但不支持 syncwarp。Weft 【67, Verification of Producer-Consumer Synchronization in GPU Programs, 2015 PLDI】处理了本文范围之外的生产者-消费者同步。Volta的数据竞争检查器优势在于其简单性和效率,它在线运行且只增加与指令数成正比的时间。
  • 代码提升(Lifting)系统:诸如C2TACO、STNG、Tenspiler和LLMLift等系统试图将底层代码提升为高级DSL。但这些系统主要针对顺序代码(C/C++/Python循环嵌套),如果应用于GPU内核,将需要额外建模共享内存、屏障和warp级同步,以及处理数据竞争和死锁的可能性。

结论

本文提出了Volta,这是首个针对机器学习GPU内核的等价性检查器。Volta直接在底层PTX汇编上操作,利用已知的张量大小和线程ID来简化必须支持的表达式空间。Volta通过解释内核以获得每个输出张量元素的符号表达式,并使用专门的规范化过程检查隐含的验证条件。其健全性和完备性由本文给出的语义以及在Agda中验证的合流性定理保证。此外,本文还证明了实数上带有指数的多项式之间恒等式的可判定性。未来的工作方向包括支持异步CUDA内置函数,如 pipelinetma