
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 python3.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 上跑用-archsm_80H100 用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, maskmask) y tl.load(y_ptr offs, maskmask) tl.store(out_ptr offs, x y, maskmask) # 获取 Triton IRTTIR/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, archsm_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_outputTrue, textTrue ) 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 参数选择的一个具体计算假设你要处理一个BLOCK128的 elementwise kernelPTX 里寄存器压力怎么估粗略算法是每个线程处理 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 和它的验证结果都存下来形成一个“生成-验证”日志。这个日志本身就是宝贵的数据——失败的样本告诉你模型的弱点在哪成功的样本可以回流做增量训练。跑上几轮你会对“模型在什么情况下会崩”有非常具体的直觉这比看任何论文都管用。