
1. 这个标题到底在说什么第一次看到“AI 就是编译器”这个说法我脑子里蹦出来的不是学术论文而是过去几年折腾 Triton 和 TVM 时那些让人头大的 lowering 报错。传统编译器后端有多难搞写过 kernel 的人都懂从高层 IR 一路降到硬件指令中间要过几十个 pass每个 pass 都可能因为一个 shape 推断失败或者 layout 冲突直接崩掉。而这篇论文的核心主张非常激进——既然大语言模型已经能理解语义和硬件约束为什么不干脆让它直接输出 PTX把整个后端链路砍掉PTX 是 NVIDIA GPU 的并行线程执行中间表示可以理解成“GPU 的汇编之上、CUDA C 之下”的一层。它比 SASS 好读比 CUDA C 更贴近硬件寄存器、共享内存、线程束这些概念都是显式暴露的。正常情况下我们写 CUDA 或者 Triton编译器负责把高层代码翻译成 PTX再由 ptxas 汇编成 SASS。这篇论文想做的事就是让 LLM 跳过前面所有环节直接吐出一段能跑的 PTX。这件事为什么值得关注因为编译器后端是出了名的“脏活累活”。Triton 之所以好用是因为它把很多 lowering 细节藏起来了但一旦你的算子超出它支持的 pattern就得回去手写 CUDA甚至手写 PTX。而手写 PTX 的门槛极高寄存器分配、bank conflict、warp divergence 这些坑没踩过几个月根本摸不清。如果 LLM 能稳定生成正确的 PTX那相当于把最稀缺的那部分能力给自动化了。这篇文章适合谁看如果你是在做算子优化、推理加速、或者对 LLM 在系统层的应用感兴趣那这篇解读值得花时间。如果你只是调调 API 写业务代码可能感受没那么深但理解这个思路对判断未来工具链的走向有帮助。我下面会从设计动机、核心技术点、实操复现、以及踩坑经验几个角度把这个工作拆开讲。2. 为什么有人想绕开编译器后端2.1 传统编译器后端的痛点在哪要理解这个工作的价值得先知道传统后端到底难在哪。以 Triton 为例你写一个tl.load和tl.store背后发生的事情是Triton IR 先做 layout 推断然后经过 coalescing pass、shared memory allocation、pipeline scheduling最后降到 LLVM IR再降到 PTX。每一步都有大量的启发式规则而这些规则是为通用场景设计的遇到特殊算子就容易失效。我印象很深的一次是做一个 fused attention 变体Triton 的自动 pipeline 把 shared memory 用爆了报了一个out of resource: shared memory的错。查了半天发现是某个 pass 的 buffer 复用策略没覆盖到我的访问模式。这种问题你没法通过调参解决只能改编译器或者换手写。而手写 CUDA 又要重新处理所有底层细节等于把编译器帮你做的事又做了一遍。TVM 那边也类似。AutoTVM 和 Ansor 能搜出不错的 schedule但 search space 的设计本身依赖人工定义的 schedule primitive。一旦算子结构变了primitive 不够用搜索就退化成随机试。而且整个 lowering 链路非常长debug 的时候要在几十个 IR 之间来回跳心智负担极大。2.2 LLM 直接生成 PTX 的逻辑论文的思路是LLM 在预训练阶段见过大量 CUDA 和 PTX 代码它对“这段计算该怎么映射到线程和寄存器”这件事已经有隐式理解。与其让它生成 CUDA 再走一遍编译器不如直接让它生成 PTX然后用 ptxas 汇编验证。这样做的优势是链路短、可控性强而且 LLM 可以根据自然语言描述里的约束比如“用 128 个线程”“避免 bank conflict”直接调整输出。这个思路和“LLM 就是编译器”的说法是吻合的把自然语言或者高层算子描述当作源语言把 PTX 当作目标语言LLM 充当翻译器。传统编译器的前端词法、语法、语义分析被 LLM 的理解能力替代后端优化、代码生成被 LLM 的生成能力替代。中间那些 pass 和 IR 全部消失。当然这不意味着传统编译器就没用了。LLM 生成的 PTX 仍然需要 ptxas 做最终汇编和寄存器分配也需要验证工具检查正确性。但至少从“高层语义到 PTX”这一段LLM 可以承担主要工作。2.3 和 Triton、TVM 的定位差异这里要澄清一个容易混淆的点这个工作不是要取代 Triton 或 TVM而是提供另一条路径。Triton 的优势在于它的编程模型对开发者友好你写的是 block-level 的代码不用管线程级细节。TVM 的优势在于它有完整的图优化和自动调优能力。而 LLM 直接生成 PTX 的优势在于灵活性和覆盖范围——只要你能用自然语言描述清楚它就能尝试生成不受限于预定义的 primitive 集合。代价也很明显LLM 生成的 PTX 正确性没有保证需要额外的验证机制。而且 PTX 本身是硬件相关的换一代 GPU 架构可能就要重新生成。所以这个方案更适合作为传统编译器的补充而不是完全替代。3. 核心技术点拆解3.1 PTX 作为目标语言的特点PTX 是 NVIDIA 定义的一种虚拟指令集它有几个特点让它适合作为 LLM 的生成目标。第一它是文本格式LLM 处理起来很自然不像 SASS 那样是二进制。第二它的抽象层级适中既有.reg这样的虚拟寄存器声明又有ld.shared、bar.sync这样的硬件相关指令LLM 可以在上面做优化决策。第三它有官方文档和大量开源代码作为训练语料LLM 对它的语法和常见 pattern 比较熟悉。但 PTX 也有坑。它的版本和 GPU 架构强相关比如cp.async在 sm_80 才引入wgmma在 sm_90 才有。LLM 如果生成了一段用了新指令的 PTX但目标 GPU 不支持ptxas 就会报错。所以实际使用中需要把目标架构作为约束传给 LLM。3.2 LLM 生成 PTX 的 prompt 设计论文里比较关键的一部分是 prompt 工程。要让 LLM 生成可用的 PTX光说“写一个矩阵乘法”是不够的。需要把以下信息都塞进 prompt算子的数学定义、输入输出的 shape 和 dtype、目标 GPU 架构、线程块大小、每个线程负责的计算量、以及任何特殊的优化要求比如用 shared memory 做 tiling。我自己的经验是prompt 里最好给一个类似的 PTX 示例让 LLM 模仿结构。比如你要生成一个 tiled matmul可以先给一段简单的 vector add 的 PTX让它理解基本的 kernel 结构然后再描述 matmul 的具体需求。这样生成的成功率会高很多。另外把约束写成“必须”和“禁止”的形式比写成建议更有效。比如“必须使用 256 个线程”“禁止使用全局内存原子操作”LLM 会更严格地遵守。3.3 正确性验证与反馈循环LLM 生成的 PTX 不能直接信必须验证。论文里用的方法是先用 ptxas 汇编如果汇编失败把错误信息反馈给 LLM 让它修正如果汇编成功再用一个参考实现比如 NumPy 或 PyTorch对比数值结果。数值对比要设置合理的容差因为浮点运算的顺序不同会导致微小差异。这个反馈循环是整套方案能跑通的关键。我试过让 LLM 一次性生成正确的 PTX成功率大概只有三成。但加上“汇编报错就重试”的循环成功率能提到七成以上。如果再加入数值验证的反馈最终可用率能到八成左右。剩下的两成通常是算子太复杂或者 LLM 对某个硬件细节理解有误。3.4 性能调优的空间生成正确的 PTX 只是第一步性能好不好是另一回事。LLM 生成的 PTX 往往能用但离手写优化过的版本还有差距。常见的性能问题包括寄存器用太多导致 occupancy 低、shared memory 访问有 bank conflict、没有用 async copy 做流水线。要提升性能可以在 prompt 里加入性能约束比如“寄存器数量控制在 64 个以内”“shared memory 访问步长设为 1 避免 bank conflict”。也可以在生成后做自动调优把线程块大小、tile 大小这些参数作为搜索维度让 LLM 生成多个版本然后实测选最优。4. 实操复现从零生成一个 PTX kernel4.1 环境准备要复现这个工作你需要以下环境一台有 NVIDIA GPU 的机器sm_70 以上比较稳妥、CUDA Toolkit包含 ptxas 和 nvcc、Python 环境用来调 LLM API 和做数值验证、以及一个能访问的 LLM本地部署或 API 都行。我用的组合是 CUDA 12.1 Python 3.10 PyTorch 2.1。LLM 方面如果本地资源够可以用 CodeLlama 或 DeepSeek-Coder 这类对代码友好的模型如果走 APIGPT-4 或 Claude 在 PTX 生成上的表现也不错。注意 PTX 的语法比较特殊通用聊天模型可能不如代码专用模型。4.2 构造 prompt 的模板下面是我实际用的 prompt 模板你可以直接抄你是一个 PTX 代码生成器。请根据以下描述生成一个完整的 PTX kernel。 算子描述{operator_description} 输入 shape{input_shapes} 输出 shape{output_shapes} 数据类型{dtype} 目标架构{target_arch} 线程块大小{block_size} 特殊要求{constraints} 要求 1. 生成的 PTX 必须能被 ptxas 汇编通过。 2. 必须使用 .visible .entry 定义 kernel。 3. 必须正确声明所有寄存器和共享内存。 4. 不要生成注释以外的任何非 PTX 内容。 参考示例 {example_ptx}这个模板的关键是把所有约束都显式写出来并且给一个示例让 LLM 模仿。示例不需要和目标任务一样但结构要相似。4.3 一个具体的生成案例我拿 vector add 做例子因为它的 PTX 结构简单适合验证流程。算子描述是“C[i] A[i] B[i]”输入是三个长度为 1024 的 float 数组线程块大小 256。LLM 生成的 PTX 大致长这样.version 7.0 .target sm_70 .address_size 64 .visible .entry vector_add( .param .u64 A, .param .u64 B, .param .u64 C, .param .u32 N ) { .reg .f32 %f4; .reg .b32 %r4; .reg .b64 %rd8; ld.param.u64 %rd1, [A]; ld.param.u64 %rd2, [B]; ld.param.u64 %rd3, [C]; ld.param.u32 %r1, [N]; mov.u32 %r2, %ctaid.x; mov.u32 %r3, %ntid.x; mov.u32 %r4, %tid.x; mad.lo.u32 %r5, %r2, %r3, %r4; setp.ge.u32 %p1, %r5, %r1; %p1 bra DONE; mul.wide.u32 %rd4, %r5, 4; add.s64 %rd5, %rd1, %rd4; add.s64 %rd6, %rd2, %rd4; add.s64 %rd7, %rd3, %rd4; ld.global.f32 %f1, [%rd5]; ld.global.f32 %f2, [%rd6]; add.f32 %f3, %f1, %f2; st.global.f32 [%rd7], %f3; DONE: ret; }这段代码能汇编通过数值结果也正确。但注意它没有做任何优化每个线程只处理一个元素没有用 vectorized load。如果要提升性能可以在 prompt 里要求“每个线程处理 4 个元素使用 ld.global.v4.f32”。4.4 验证脚本的写法验证分两步先汇编再跑数值对比。汇编用ptxas -archsm_70 kernel.ptx -o kernel.cubin如果报错就把错误信息拼回 prompt 让 LLM 重试。数值对比用 PyTorch 加载 cubin 比较麻烦简单做法是用 ctypes 调 CUDA driver API或者用 cupy 的 RawKernel。我一般用 cupy因为它对 PTX 的支持比较直接import cupy as cp ptx_code open(kernel.ptx).read() module cp.RawModule(codeptx_code, backendnvrtc) kernel module.get_function(vector_add) A cp.random.randn(1024, dtypecp.float32) B cp.random.randn(1024, dtypecp.float32) C cp.zeros(1024, dtypecp.float32) kernel((4,), (256,), (A, B, C, 1024)) assert cp.allclose(C, A B, atol1e-5)注意backendnvrtc这里cupy 默认用 nvrtc 编译如果你的 PTX 版本和 nvrtc 不兼容可以改成用 driver API 加载 cubin。5. 常见问题与排查技巧5.1 汇编报错怎么定位最常见的汇编错误是寄存器声明不匹配。比如你用了%f1但没在.reg .f32里声明ptxas 会报Undeclared identifier。另一个常见错误是类型不匹配比如把.f32的值存到.b32的寄存器里。排查方法是看 ptxas 报的行号然后对照 PTX ISA 文档检查那一条指令的操作数类型。还有一种错误是版本不兼容。比如你用了cp.async但.target写的是sm_70ptxas 会报Feature not supported on target。这时候要么改 target要么换指令。5.2 数值不对怎么查数值不对通常有几个原因索引算错了、边界条件没处理、或者浮点累加顺序不同导致精度差异。排查方法是先用小规模输入比如 N16手动算一遍对比 PTX 的输出。如果小规模对但大规模不对多半是边界条件的问题检查setp和bra的逻辑。浮点精度问题比较难查因为 PTX 里的add.f32和 PyTorch 的加法在数学上等价但实际执行顺序可能不同。如果差异在 1e-5 以内一般可以接受如果差异很大那说明逻辑有问题。5.3 性能不达预期怎么办LLM 生成的 PTX 性能差是常态因为它在生成时没有做 cost model 评估。提升性能的入手点有几个第一检查 global memory 访问是否合并如果每个线程访问的地址不连续带宽会浪费。第二检查 shared memory 是否有 bank conflict步长设为 1 通常能避免。第三检查寄存器数量用ptxas -v可以看到寄存器和 shared memory 的使用量如果寄存器超过 128 个occupancy 会很低。如果手动调优太慢可以写一个简单的搜索脚本把 block size、tile size、unroll factor 作为参数让 LLM 生成多个版本然后实测选最快的。我试过在 matmul 上做这个搜索大概 20 个版本里能找到一个接近手写性能的。5.4 常见问题速查表问题现象可能原因排查方法ptxas 报 Undeclared identifier寄存器未声明检查 .reg 声明是否覆盖所有使用的寄存器ptxas 报 Feature not supported指令与 target 不匹配检查 .target 版本和指令引入版本数值结果全错索引或边界逻辑错误用小规模输入手动验证数值结果有微小差异浮点累加顺序不同检查容差设置确认是否在可接受范围性能远低于预期内存访问未合并或 occupancy 低用 ptxas -v 看资源使用检查访问模式LLM 生成的 PTX 不完整prompt 约束不够在 prompt 里明确要求完整 kernel 结构6. 这个方向后续能怎么扩展6.1 从 PTX 到多后端支持目前这个工作聚焦在 PTX但同样的思路可以扩展到其他后端。比如让 LLM 直接生成 SPIR-VVulkan 的计算着色器中间表示或者生成 AMD 的 GCN 汇编。核心逻辑是一样的把高层描述当作源语言把硬件中间表示当作目标语言LLM 充当翻译器。不同后端的差异在于 prompt 里需要注入的硬件约束不同。6.2 和自动调优结合LLM 生成 PTX 的一个弱点是性能不稳定而自动调优正好能补上这块。可以把 LLM 当作 schedule 生成器让它根据算子描述生成多个候选 PTX然后用实测数据做反馈逐步优化 prompt。这个循环跑几轮之后生成质量会明显提升。6.3 在推理框架里的落地场景最直接的落地场景是自定义算子的快速原型。比如你在做某个新模型的推理加速遇到一个 PyTorch 没有优化实现的算子传统做法是手写 CUDA 或者等 Triton 支持。用 LLM 生成 PTX 可以在几分钟内得到一个能跑的版本虽然性能不一定最优但至少能先验证正确性然后再决定要不要深度优化。另一个场景是教学和调试。PTX 可读性比 SASS 好LLM 生成的 PTX 可以作为学习材料帮助理解高层代码到硬件的映射关系。我在带新人的时候会让他们对比 LLM 生成的 PTX 和 nvcc 生成的 PTX看看差异在哪这对理解编译器优化很有帮助。6.4 当前的局限和注意事项这个方案目前有几个硬伤。第一LLM 对 PTX 的理解深度有限复杂算子比如涉及 wgmma 或 TMA 的生成成功率很低。第二PTX 和 GPU 架构强绑定换架构就要重新生成维护成本不低。第三正确性验证依赖参考实现如果参考实现本身有问题验证就失效了。所以我的建议是把 LLM 生成的 PTX 当作草稿不要直接上生产。用它来加速原型验证和教学可以但关键路径上的算子还是要有手写优化版本兜底。另外prompt 里的约束要写得非常具体模糊的描述会导致生成结果不可控。我个人在实际操作中的体会是这套方法最适合的场景是“中等复杂度、有明确参考实现、性能要求不极端”的算子。太简单的算子用 Triton 就够了太复杂的算子 LLM 搞不定。中间那部分——Triton 支持不好但手写又太费时的——才是 LLM 生成 PTX 的甜区。