1. 这篇论文到底想干什么:把编译器后端整个拿掉
第一次看到“AI 就是编译器”这个说法,我脑子里蹦出来的画面是:一个模型坐在原本属于 LLVM 后端的位置上,输入是高层中间表示,输出直接就是能在 GPU 上跑的 PTX 汇编。这个想法乍一听有点离谱,但仔细想想,它戳中的恰恰是编译器工程里最贵、最脏、最难维护的那一段——lowering 和指令选择。
传统编译流程里,从 Triton、TVM 或者各种 DSL 到最终 GPU 可执行代码,中间要经过好几层:高层 IR 优化、循环变换、tile 化、向量化、寄存器分配、指令调度,最后才落到 PTX 或者 SASS。每一层都有大量手写规则、pattern matching、启发式代价模型。这套东西能跑,但极其脆弱:换一代硬件、换一种算子形态、换一个数据布局,后端工程师就得重新调一遍。论文的核心主张就是——既然大模型已经能理解高层语义和硬件约束,为什么不让它直接干“高层 IR 到 PTX”这一步,把整个后端当成一个被学出来的函数?
我先把结论摆前面:这篇工作不是要证明“LLM 能替代编译器”,而是想验证一个更窄但更关键的命题——在受限的算子集合和固定的目标架构下,LLM 生成的 PTX 在正确性和性能上能不能逼近甚至超过手写后端。这个定位很重要,因为它决定了你看这篇论文时该关注什么:不是“通用性”,而是“在特定 lowering 任务上,学习式方法是否已经具备工程可用性”。
适合读这篇内容的人,我大致分三类。第一类是做 AI 编译栈的工程师,尤其是天天跟 Triton、TVM、MLIR 打交道、被后端 bug 折磨过的人;第二类是做 LLM for code 的研究者,想看看代码生成从 Python、C++ 往汇编级别下沉会遇到什么新问题;第三类是想理解“AI 与系统软件结合”这条路线到底走到哪一步的技术管理者。如果你只是想知道“怎么用 LLM 写个排序算法”,这篇不适合你。
关键词里出现的PTX、LLM、Triton、编译器后端、AI lowering,基本就是全文的骨架。我下面会按“为什么这么设计 → 核心机制怎么拆 → 实操上怎么复现 → 会踩哪些坑”这个顺序展开,尽量把论文里没写透、但工程上必须知道的细节补上。
2. 为什么盯上 PTX:选型背后的真实考量
2.1 PTX 是“可读的硬件契约”,不是随便挑的
很多人第一反应是:为什么不直接生成 SASS(真正的机器码)?答案很现实——PTX 是 NVIDIA 官方文档化、稳定、跨代兼容的中间汇编,而 SASS 是未公开的、跟具体 SM 架构强绑定的。让 LLM 去学生成 SASS,等于让它去拟合一个没有规范、随时会变的黑盒,训练信号和验证都无从下手。PTX 则不同,它有完整的 ISA 手册,指令语义清晰,寄存器模型、内存空间、barrier、warp shuffle 这些都有明确定义。
更关键的是,PTX 可以被ptxas汇编、被nvdisasm反汇编、被cuobjdump检查,也就是说验证链路是现成的。你生成一段 PTX,能不能编译、编译出来对不对、跑起来性能如何,全都有工具兜底。这一点对“学习式 lowering”是生死攸关的:没有可靠的自动验证,整个方法就没法闭环。
2.2 绕开后端,本质是把“规则工程”换成“数据工程”
传统后端的工作量有多大?我举个具体例子。一个tl.dot在 Triton 里,要经过 layout 推导、shared memory 分配、mma 指令选择、pipeline 调度,最后生成几十到上百条 PTX。这些逻辑是工程师一条条写死的,覆盖的是“已知的算子 + 已知的 shape + 已知的架构”。一旦出现新的 fused pattern,比如 attention 里那种带 mask、带 scale、带 softmax 的复合结构,后端要么写新 pass,要么退化成低效的通用路径。
论文的思路是:把这些规则从代码里搬到模型权重里。训练数据就是“高层 IR 片段 → 对应的高质量 PTX”这样的配对。模型学到的不是某条规则,而是“给定这段计算意图和这些硬件约束,PTX 大概长什么样”的分布。这样做的好处是,面对训练分布内的新组合,模型有可能直接泛化出合理代码,而不需要工程师再写一条 pattern。
但代价也很明显:可解释性和可调试性下降。手写后端出 bug,你可以定位到某个 pass;模型生成的 PTX 出 bug,你只能看输入输出,中间是黑盒。所以论文在评估里特别强调正确性验证,这不是走过场,而是这个方法能不能被信任的前提。
2.3 和 Triton 的关系:不是替代,是补位
这里要澄清一个常见误解。Triton 本身已经是一个“高层 DSL + 编译器”的方案,它的后端也是基于 MLIR 和 LLVM 的。论文并不是说“Triton 没用了”,而是说Triton 到 PTX 这一段 lowering,可以尝试用 LLM 来做。实际上,Triton 的 IR 非常适合作为 LLM 的输入:它比 CUDA C 更抽象,去掉了大量语法噪音,又比纯数学表达式更接近硬件,保留了 block、thread、shared memory 这些概念。
我在实际看这类工作时的一个判断标准是:输入表示是否“语义密度高且噪音低”。Triton IR 恰好满足。如果输入是原始 Python 或者带大量模板的 CUDA,模型要花很多容量去理解语法;如果输入是纯数据流图,又丢失了并行结构信息。Triton 卡在中间,这是它被选中的深层原因。
3. 核心机制拆解:LLM 到底怎么“当编译器”
3.1 任务形式化:从序列到序列,但不止是翻译
表面上看,这就是个 seq2seq 任务:输入 Triton IR 文本,输出 PTX 文本。但如果你真按机器翻译那套去做,基本会失败。原因是 PTX 有强结构约束:寄存器必须先声明后使用,barrier 必须成对出现,shared memory 访问要符合对齐要求。这些约束不是统计规律,而是硬性规则。
论文采用的做法,我理解是约束解码 + 后验验证的组合。约束解码保证生成的 token 序列在语法上合法,比如寄存器命名、指令格式;后验验证则是把生成的 PTX 丢给ptxas编译,编译不过就丢弃或重采样。这个“生成-验证-筛选”的循环,是让 LLM 输出从“看起来像”变成“真的能用”的关键。
3.2 训练数据的构造:质量比数量重要得多
这类工作最容易被低估的就是数据。你不能随便抓一堆 Triton 代码和对应的 PTX 就开训,因为编译器生成的 PTX 质量参差不齐,而且同一个 IR 在不同优化级别下 PTX 差异巨大。论文里大概率做了这几件事:
- 固定编译配置:统一优化级别、统一目标架构(比如 sm_80 或 sm_90),消除配置带来的噪声。
- 筛选高质量样本:只保留性能达标、无冗余指令的 PTX,可能用
ptxas -v的寄存器占用和指令数做过滤。 - 对齐粒度:不是整个 kernel 对整段 PTX,而是按基本块或按算子切分,降低单样本复杂度。
我自己的经验是,数据对齐粒度决定了模型能学到什么。如果按整个 kernel 训,模型学到的是“整体结构”;如果按基本块训,学到的是“局部指令选择”。论文如果同时用了两种粒度,那说明它在兼顾全局调度和局部 lowering。
3.3 推理时的关键:怎么保证生成的 PTX 真的对
这是整个方法最脆弱也最核心的环节。我总结下来有三道关:
第一道是语法关,靠约束解码或者语法引导的 beam search,保证生成的 PTX 能被 parser 接受。第二道是编译关,用ptxas实际汇编,失败就重试。第三道是数值关,把编译出的 cubin 加载运行,和参考实现对比输出,误差超过阈值就判定失败。
这三道关里,第二道是成本最低、过滤效果最好的。很多语法合法的 PTX 其实过不了ptxas,比如寄存器类型不匹配、shared memory 超限。第三道最贵,但只有它能抓住“编译通过但算错”的情况,比如 race condition、精度问题。论文如果报告了端到端正确率,那一定是三道关都过了的比例,这个数字通常比纯语法正确率低不少,但才是真正有意义的指标。
4. 实操复现:如果你想自己跑一遍
4.1 环境准备与依赖
要复现这类工作,硬件上你需要一块 NVIDIA GPU,架构最好和论文一致(比如 A100 对应 sm_80)。软件栈大致是:
# 基础环境 conda create -n ai-compiler python=3.10 conda activate ai-compiler # Triton(用于生成 IR 和参考 PTX) pip install triton # CUDA Toolkit(提供 ptxas、nvdisasm) # 确保 nvcc 和 ptxas 在 PATH 里 nvcc --version ptxas --version # 训练框架 pip install torch transformers datasets accelerate提示:
ptxas的版本要和目标架构匹配。用ptxas --help可以看到支持的-arch选项。如果你在 A100 上跑,用-arch=sm_80;H100 用sm_90。版本不匹配会导致明明正确的 PTX 编译失败。
4.2 构造训练样本的脚本思路
我写过一个简化版的样本构造流程,核心是“用 Triton 生成 IR,用官方后端生成 PTX,配对保存”:
import triton import triton.language as tl import subprocess import tempfile import os @triton.jit def add_kernel(x_ptr, y_ptr, out_ptr, n, BLOCK: tl.constexpr): pid = tl.program_id(0) offs = pid * BLOCK + tl.arange(0, BLOCK) mask = offs < n x = tl.load(x_ptr + offs, mask=mask) y = tl.load(y_ptr + offs, mask=mask) tl.store(out_ptr + offs, x + y, mask=mask) # 获取 Triton IR(TTIR/TTGIR) ir = add_kernel.warmup(...).asm["ttgir"] # 获取 PTX ptx = add_kernel.warmup(...).asm["ptx"] # 保存配对 with open("sample_0001.ir", "w") as f: f.write(ir) with open("sample_0001.ptx", "w") as f: f.write(ptx)这里有个细节:warmup需要提供具体的参数和 grid,否则拿不到编译结果。我一般会写一个小工具函数,把 shape、dtype、grid 都参数化,批量生成不同配置的样本。样本多样性主要来自 shape、block size、是否带 mask、是否带 reduction,这些维度覆盖得越全,模型泛化越好。
4.3 验证流水线的搭建
生成归生成,验证才是重头戏。我建议把验证做成一个独立模块,输入是 PTX 字符串,输出是“是否可用 + 性能数据”:
def validate_ptx(ptx_str, ref_fn, inputs, arch="sm_80"): with tempfile.TemporaryDirectory() as d: ptx_path = os.path.join(d, "k.ptx") cubin_path = os.path.join(d, "k.cubin") with open(ptx_path, "w") as f: f.write(ptx_str) # 第一步:编译 r = subprocess.run( ["ptxas", f"-arch={arch}", ptx_path, "-o", cubin_path], capture_output=True, text=True ) if r.returncode != 0: return {"ok": False, "stage": "compile", "err": r.stderr} # 第二步:加载运行,对比数值 # 这里用 cuda-python 或 pycuda 加载 cubin # 第三步:计时 return {"ok": True, "stage": "run"}注意:
ptxas编译通过不代表能加载。有时候 PTX 里用了目标架构不支持的指令,编译会过但加载失败。所以第二步的加载测试不能省。
4.4 参数选择的一个具体计算
假设你要处理一个BLOCK=128的 elementwise kernel,PTX 里寄存器压力怎么估?粗略算法是:每个线程处理 1 个元素,需要至少 2 个输入寄存器 + 1 个输出寄存器 + 若干地址寄存器,大约 8-12 个寄存器。如果模型生成的 PTX 用了 40 个寄存器,那 occupancy 就会明显下降。我在筛选样本时会把ptxas -v输出的寄存器数作为过滤条件,超过阈值(比如 32)的样本直接丢掉,避免模型学到“浪费寄存器”的坏习惯。
5. 常见问题与排查技巧实录
5.1 生成结果速查表
| 现象 | 可能原因 | 排查手段 | 解决方向 |
|---|---|---|---|
| ptxas 报语法错误 | 寄存器未声明、指令格式错 | 看 stderr 行号 | 加强约束解码 |
| 编译过但加载失败 | 用了不支持的指令 | cuobjdump 看指令 | 限制指令集范围 |
| 能跑但结果错 | race condition、精度 | 对比参考输出 | 加 barrier 约束 |
| 结果对但很慢 | 寄存器过多、无向量化 | ptxas -v 看占用 | 性能过滤样本 |
| 换个 shape 就崩 | 过拟合到固定配置 | 测不同 shape | 增加数据多样性 |
5.2 我踩过的几个坑
第一个坑是把 PTX 当纯文本处理。早期我直接用 tokenizer 切 PTX,结果寄存器名%r1和%r10被切成完全不同的 token,模型学不到“寄存器编号是连续的”这个规律。后来改成按 PTX 语法做 tokenization,把%r和数字分开,效果明显好转。
第二个坑是忽略编译配置。同一段 IR,开-O3和不开,PTX 差很多。如果训练数据混了不同优化级别,模型会学乱。我的做法是全部固定-O3,并且在输入里显式带上架构信息,让模型知道目标是什么。
第三个坑是验证不充分。有次模型生成的 PTX 在小 shape 上全对,我一度以为成了,结果一上大 shape 就出 race condition。后来我把验证集按 shape 分层,小、中、大各占三分之一,才暴露出来。数值验证一定要覆盖边界 shape,这是血泪教训。
5.3 性能对比的注意事项
论文里如果报告了“超过手写后端”,你要特别小心看它的对比基线。常见的情况是:基线用的是未调优的通用路径,而模型生成的是针对特定 shape 特化的代码。这种对比不公平。我建议自己复现时,基线一定要用同一套 Triton 配置、同一优化级别生成的 PTX,这样比出来的差距才是方法本身的差距。
另外,性能测量要用 CUDA event 而不是 CPU 计时,要 warmup 足够次数,要排除首次加载的开销。这些是 GPU 性能测试的基本功,但很多人图省事就忽略了,导致数据不可信。
6. 这条路线的边界与我的判断
6.1 它现在能做什么,不能做什么
从论文的定位和现有结果看,在固定架构、固定算子族、固定 shape 分布内,LLM 直接生成 PTX 是可行的,正确率能做到可用水平,性能能接近甚至局部超过手写后端。但出了这个范围,比如换架构、换全新算子、遇到极端 shape,可靠性会快速下降。
这不是方法本身的缺陷,而是所有学习式方法的共性。它的价值不在于“通用”,而在于把后端工程师从重复的 pattern 编写中解放出来。你可以想象一个工作流:新算子先让模型生成一版 PTX,工程师 review 和微调,再固化成规则。这样人力集中在真正新的问题上,而不是重复劳动。
6.2 对 AI lowering 这个方向的看法
“AI lowering”这个词最近出现频率很高,但我觉得要区分两种含义。一种是用 AI 辅助编译器做决策,比如用模型预测 tile size、预测是否向量化,这是增强现有编译器。另一种是用 AI 替代编译器的某个阶段,也就是这篇论文做的,直接生成目标代码。前者风险低、易落地,后者激进但天花板高。
我的判断是,短期内前者会更早进入生产环境,因为它可以嵌在现有流程里,出错了有兜底。后者更适合作为研究探索,积累数据和经验。但长期看,如果模型对硬件的理解足够深,后者有可能反过来重塑编译器的架构——后端不再是手写 pass 的集合,而是“模型 + 验证器 + 少量规则”的组合。
6.3 给想跟进的人的建议
如果你打算在这个方向做点东西,我的建议是先把验证基础设施做扎实。很多人一上来就调模型,结果生成的东西对不对都判断不了,纯属浪费时间。先把ptxas编译、cubin 加载、数值对比、性能计时这条链路跑通,再去做生成。另外,从小算子开始,比如 elementwise、reduction,别一上来就搞 attention,那个复杂度会让你怀疑人生。
数据方面,宁缺毋滥。一百条高质量、配置统一的样本,比一万条混杂的样本有用得多。模型方面,不用追求最大,7B 级别的代码模型在充分微调后,在受限任务上表现已经可以接受。关键是任务定义要窄、验证要严、迭代要快。
最后分享一个我在实际搭建这类流水线时的小技巧:把每次生成的 PTX 和它的验证结果都存下来,形成一个“生成-验证”日志。这个日志本身就是宝贵的数据——失败的样本告诉你模型的弱点在哪,成功的样本可以回流做增量训练。跑上几轮,你会对“模型在什么情况下会崩”有非常具体的直觉,这比看任何论文都管用。