☰
LLM 直接生成 PTX:绕过编译器后端的 GPU 编程新思路
2026/10/8 1:17:05 网站建设 项目流程

1. 这篇论文到底在讲什么

第一次看到“AI 就是编译器”这个说法,我的反应是:又是一个标题党。但把论文翻完之后,我改主意了——它讨论的问题非常具体,而且戳中了当下 GPU 编程工具链里一个真实存在的痛点。

先把背景交代清楚。我们平时写 GPU 代码,主流路径大概是这样:写 CUDA C++ 或者 Triton,然后交给编译器(NVCC、LLVM、Triton 自己的后端)一路 lowering,最终生成 PTX(Parallel Thread Execution),再由 ptxas 汇编成 SASS 在显卡上跑。这条链路成熟、稳定,但层级多、抽象厚。你想精确控制某个 warp 的寄存器分配、想手动安排 shared memory 的 bank 布局,往往要跟编译器“斗智斗勇”,写一堆 pragma 还不一定听话。

这篇论文提出的思路是:既然大语言模型已经能写 CUDA、能写 Triton,那能不能让它跳过中间所有抽象层,直接输出 PTX?换句话说,把 LLM 当成一个“编译器后端”,输入是自然语言描述或者高层代码意图,输出是可直接被 ptxas 接受的 PTX 汇编。

这个想法乍一听很疯狂,因为 PTX 是接近硬件的低级 IR,寄存器、线程索引、内存空间限定符一个都不能错。但论文的核心论点恰恰是:LLM 在预训练阶段已经“见过”海量 PTX 代码(CUDA 工具链、开源项目、反汇编数据里都有),它对 PTX 的语法和常见模式其实有相当强的先验。与其让它生成高层代码再走一遍可能“优化过头”或“优化不到位”的后端,不如让它直接产出目标汇编。

我个人的判断是:这篇论文的价值不在于“以后不用编译器了”,而在于它打开了一个新的视角——LLM 可以作为一种可编程的 lowering 引擎。传统编译器后端是确定性的、基于规则的,而 LLM 后端是概率性的、基于模式的。两者各有适用场景,前者追求正确性和可复现,后者追求灵活性和“意图对齐”。

适合读这篇内容的人:做 GPU kernel 优化的工程师、研究 AI for systems 的同学、以及想搞清楚 LLM 到底能不能碰底层代码的开发者。如果你只是调 PyTorch 的 API,这篇可能离你有点远,但了解这个方向对判断未来工具链走向有帮助。

2. 为什么有人想绕开编译器后端

2.1 传统编译链路的“抽象税”

要理解这篇论文的动机,得先明白传统链路哪里让人不爽。

从 CUDA C++ 到 SASS,中间至少经过这几层:CUDA C++ → LLVM IR(NVVM)→ PTX → SASS。每一层都有自己的优化 pass,每一层都可能做出跟你预期不一样的决定。比如你写了一个循环,想让编译器展开成 4 路,结果它展开了 8 路,寄存器压力爆了,occupancy 掉了一半。你写__shared__数组想手动控制 bank conflict,结果编译器给你重排了访问顺序。

这就是所谓的“抽象税”——你为了用高级语言,付出了对底层失去精确控制的代价。对于 90% 的场景,这个税是值得交的,因为编译器比你聪明。但对于那 10% 的极致优化场景(比如 flash attention 的早期手写 kernel、某些量化算子的融合),工程师往往要反汇编看 SASS,然后回头改 CUDA 代码“诱导”编译器生成想要的指令。这个过程极其痛苦。

Triton 的出现缓解了一部分问题,它把抽象层级降到了 block 级别,让你用 Python 语法描述 tile 操作。但 Triton 依然有自己的后端,依然会做它认为合理的优化。你想控制的东西,它不一定暴露给你。

2.2 LLM 作为 lowering 引擎的合理性

论文的切入点就在这里:如果 LLM 已经学会了 PTX 的“语法 + 常见模式”,那它能不能直接根据意图生成 PTX,跳过中间所有可能“误解你”的层?

这个想法的合理性来自几个观察:

第一,PTX 虽然低级,但它的指令集是有限的、模式化的。常见的操作就那么几十种:ld.global、st.shared、mad.f32、bar.sync、shfl.sync等等。LLM 在预训练中见过大量 PTX 片段(尤其是开源 CUDA 项目编译后的产物),对这些指令的用法有统计意义上的“语感”。

第二,很多 kernel 的结构是高度模板化的。比如一个 reduction kernel,它的 PTX 骨架基本固定:加载、warp shuffle 归约、shared memory 跨 warp 归约、写回。LLM 完全可以学会这个模板,然后根据具体的数据类型、block size 做参数化生成。

第三,LLM 的“意图理解”能力可以弥补传统编译器在语义层面的不足。你用自然语言说“我要一个对 float16 做 warp-level reduction 的 kernel,block size 256”,LLM 能直接映射到对应的 PTX 模式,而不需要你写一堆代码再祈祷编译器优化对。

注意:这里的“绕开编译器后端”不是说要抛弃 ptxas。PTX 本身还是需要 ptxas 汇编成 SASS 的。论文绕开的是从高层语言到 PTX 之间的那些 lowering pass,而不是最后的汇编步骤。

2.3 和 Triton、TVM 等方案的对比

有人会问:这不就是 Triton 在做的事吗?不完全是。

Triton 的定位是“用 Python 写 tile 级程序”,它的后端依然是确定性的编译器。你写tl.load、tl.store,Triton 编译器负责把它 lowering 成 PTX。这个过程是规则驱动的,可复现,但灵活性受限于 Triton 暴露的 API。

TVM 走的是另一条路:用 schedule 原语描述优化,然后由 TVM 的代码生成器产出目标代码。它的抽象层级比 Triton 更低,但学习曲线陡峭,而且 schedule 空间是人为定义的。

这篇论文的方案是:把 lowering 这一步交给 LLM。好处是灵活性极高——你可以用自然语言描述任何意图,LLM 尝试生成对应的 PTX。坏处是正确性没有保证——LLM 可能生成语法正确但语义错误的 PTX,而且同一个 prompt 两次生成的结果可能不一样。

所以我的看法是:这不是替代关系,而是互补关系。Triton 适合“我要写一个标准的 matmul,帮我自动优化”,LLM-as-compiler 适合“我要一个非常规的、编译器可能优化不好的 kernel,我直接用 PTX 表达意图”。

3. 核心技术点拆解

3.1 PTX 到底长什么样

为了让后面的讨论有共同语言,先快速过一下 PTX 的基本形态。PTX 是 NVIDIA 定义的一种虚拟 ISA,它有几个关键特征:

  • 显式的线程索引:%tid.x、%ctaid.x、%ntid.x这些特殊寄存器直接暴露给你。
  • 显式的内存空间:.global、.shared、.local、.const等限定符必须写清楚。
  • 显式的寄存器声明:.reg .f32 %f<10>声明 10 个 float 寄存器。
  • 虚拟寄存器:%f1、%r2这些是虚拟的,ptxas 会做寄存器分配。

一个最简单的 vector add kernel 的 PTX 大概是这样:

.version 7.0 .target sm_80 .address_size 64 .visible .entry vec_add( .param .u64 a_ptr, .param .u64 b_ptr, .param .u64 c_ptr, .param .u32 n ) { .reg .f32 %f<4>; .reg .b32 %r<6>; .reg .b64 %rd<8>; ld.param.u64 %rd1, [a_ptr]; ld.param.u64 %rd2, [b_ptr]; ld.param.u64 %rd3, [c_ptr]; ld.param.u32 %r1, [n]; mov.u32 %r2, %ctaid.x; mov.u32 %r3, %ntid.x; mov.u32 %r4, %tid.x; mad.lo.s32 %r5, %r2, %r3, %r4; setp.ge.s32 %p1, %r5, %r1; @%p1 bra DONE; mul.wide.s32 %rd4, %r5, 4; add.s64 %rd5, %rd1, %rd4; ld.global.f32 %f1, [%rd5]; add.s64 %rd6, %rd2, %rd4; ld.global.f32 %f2, [%rd6]; add.f32 %f3, %f1, %f2; add.s64 %rd7, %rd3, %rd4; st.global.f32 [%rd7], %f3; DONE: ret; }

这段代码信息量很大。你能看到线程索引怎么算、边界检查怎么做、地址怎么从 32 位扩展到 64 位、load/store 怎么带内存空间限定符。LLM 要生成这样的代码,必须对这些模式有精确的掌握。

3.2 LLM 生成 PTX 的三种输入模式

论文里(以及我看到的类似工作)通常考虑三种输入模式,难度递增:

模式一:自然语言描述 → PTX

输入是“写一个对 float32 数组做 element-wise 加法的 kernel,block size 256”。LLM 需要理解意图,选择合适的指令,生成完整 PTX。这个模式最灵活,但正确性最难保证。

模式二:CUDA C++ → PTX

输入是一段 CUDA 代码,LLM 直接输出对应的 PTX。这相当于让 LLM 扮演 NVCC 的角色。好处是有明确的语义参照,LLM 可以“翻译”而不是“创造”。坏处是 CUDA 到 PTX 的映射有很多细节(比如__syncthreads()对应bar.sync 0),LLM 必须准确掌握。

模式三:Triton → PTX

输入是 Triton 的 tile 级代码,LLM 输出 PTX。这个模式介于前两者之间,因为 Triton 本身有明确的语义,但它的 lowering 规则比 CUDA 更复杂(涉及 tile 的 layout 转换)。

从实操角度看,模式二最容易做对,因为 CUDA 和 PTX 之间的对应关系相对直接。模式一最有想象力,但需要大量的验证和纠错机制。

3.3 正确性验证:LLM 生成的东西能跑吗

这是整个方案最关键的环节。LLM 生成的 PTX 可能有几类错误:

  • 语法错误:指令拼写错、寄存器类型不匹配、缺少必要的声明。这类错误 ptxas 会直接报错,容易发现。
  • 语义错误:语法正确但逻辑错,比如边界检查写反了、地址计算偏移错了。这类错误 ptxas 不会报,但运行结果不对。
  • 性能错误:能跑但跑得慢,比如该用 shared memory 的地方用了 global memory,该用 vectorized load 的地方用了标量 load。

论文通常采用“生成-验证-修复”的循环:LLM 生成 PTX → ptxas 编译 → 如果编译失败,把错误信息喂回 LLM 让它修复 → 如果编译成功,跑单元测试验证数值正确性 → 如果数值不对,把差异信息喂回 LLM。

这个循环的有效性取决于 LLM 的“自我纠错”能力。实测下来,语法错误通常一两轮就能修好,语义错误需要更精确的错误反馈(比如告诉它“第 3 个元素的结果应该是 X,你算出来是 Y”)。

实操心得:在验证环节,不要只跑一个测试用例。PTX 的错误往往在边界条件下才暴露,比如 n 不是 block size 整数倍的时候、数组长度为 0 的时候。我建议至少准备 5 组测试数据,覆盖正常、边界、异常三种情况。

4. 实操复现:从零搭一个 LLM-to-PTX 流程

4.1 环境准备与工具选型

如果你想自己复现这个流程,需要准备这些东西:

  • LLM:可以用 API 调用,也可以本地跑。本地跑的话,建议至少 13B 参数以上的模型,7B 的模型对 PTX 这种低资源语言掌握不够。如果要用本地模型,GGUF 格式配合 llama.cpp 是比较省事的方案,安卓上都能跑(虽然性能有限)。
  • CUDA Toolkit:需要 ptxas 来验证生成的 PTX。装 CUDA Toolkit 的时候注意版本,PTX 的版本和 sm 架构要匹配。
  • 测试框架:Python 的 pytest 或者简单的脚本都行,关键是要能自动跑数值对比。
  • 参考 PTX 库:准备一批“标准答案”PTX,用来做 few-shot 示例或者微调数据。

工具选型上,我的建议是:先用 API 跑通流程,验证可行性,再考虑本地部署。因为 PTX 生成对模型能力要求高,小模型很容易生成一堆废代码,调试成本很高。

4.2 Prompt 设计与 few-shot 策略

Prompt 的设计直接决定生成质量。我试过的几种策略:

策略一:零样本 + 详细指令

直接告诉模型“你是一个 PTX 代码生成器,根据以下描述生成 PTX”,然后附上 PTX 的语法要点。效果一般,模型容易漏掉细节。

策略二:Few-shot + 标准示例

在 prompt 里放 2-3 个完整的“描述 → PTX”示例,让模型模仿。效果明显好于零样本,尤其是示例覆盖了边界检查、地址计算这些关键模式的时候。

策略三:分步生成

先让模型生成 kernel 的伪代码或者步骤列表,再让它把每一步翻译成 PTX。这个策略的好处是模型不容易“跳步”,坏处是 token 消耗翻倍。

我实测下来,策略二性价比最高。示例的选择很关键:最好覆盖 vector add、reduction、matmul 这三种典型模式,因为大部分 kernel 都是它们的变体。

一个 few-shot 示例的骨架大概是这样:

描述:对 float32 数组做 element-wise 加法,block size 256 PTX: .version 7.0 .target sm_80 ...(完整 PTX) 描述:对 float32 数组做 block-level reduction PTX: ...(完整 PTX) 描述:{用户的新需求} PTX:

4.3 生成-编译-测试的自动化脚本

整个流程可以用一个 Python 脚本串起来。核心逻辑:

import subprocess import tempfile import os def generate_ptx(prompt, model): # 调用 LLM 生成 PTX response = model.generate(prompt) return extract_ptx(response) def compile_ptx(ptx_code, arch="sm_80"): # 写临时文件,调用 ptxas 编译 with tempfile.NamedTemporaryFile(suffix=".ptx", delete=False) as f: f.write(ptx_code.encode()) ptx_path = f.name result = subprocess.run( ["ptxas", "-arch=" + arch, ptx_path, "-o", "/dev/null"], capture_output=True, text=True ) os.unlink(ptx_path) return result.returncode == 0, result.stderr def test_correctness(ptx_code, test_cases): # 把 PTX 加载成 CUDA module,跑测试用例 # 这里需要用到 cuda-python 或者 pycuda ...

这个脚本的关键在于错误处理:如果 ptxas 报错,要把 stderr 的内容整理成人类可读的反馈,再喂回 LLM。如果数值测试失败,要计算出具体哪个元素错了、期望值是多少、实际值是多少。

4.4 一个完整的 vector add 案例

我拿 vector add 做过完整测试。输入描述是:“生成一个 PTX kernel,对两个 float32 数组做逐元素加法,结果写入第三个数组,数组长度 n 通过参数传入,block size 256,需要边界检查。”

第一次生成,模型漏了边界检查,直接无条件 load/store。跑测试的时候 n=1000 就段错误了。把错误信息(“访问越界,n=1000 时第 1000 个线程访问了非法地址”)喂回去,第二次生成加上了setp.ge和@%p1 bra。再跑,通过。

第三次我让它“优化一下,用 vectorized load”,它把ld.global.f32改成了ld.global.v4.f32,但地址计算没改,导致对齐错误。ptxas 没报错,但运行结果不对。把“结果第 4 个元素开始全错”喂回去,它修正了地址步长。

这个过程大概迭代了 5 轮,最终生成的 PTX 能正确运行,性能大概是手写 CUDA 编译后的 80% 左右。差距主要在寄存器分配和指令调度上,LLM 生成的 PTX 比较“直白”,没有做深度优化。

5. 常见问题与排查技巧

5.1 生成结果不稳定怎么办

同一个 prompt,两次生成可能完全不同。这是概率模型的固有特性。缓解方法:

  • 降低 temperature:把采样温度调到 0.2 以下,生成会更确定,但可能陷入局部最优。
  • 多次生成 + 筛选:生成 5 个候选,选第一个能通过编译和测试的。
  • 固定随机种子:如果 API 支持,固定 seed 可以复现结果。

5.2 ptxas 报错看不懂怎么反馈给 LLM

ptxas 的错误信息有时候很晦涩,比如“Arguments mismatch for instruction 'ld'”。直接把这句喂给 LLM 效果不好。我的做法是:把出错的那一行 PTX 也附上,再加上前后各两行上下文,让 LLM 自己定位问题。

5.3 数值正确但性能差怎么优化

LLM 生成的 PTX 通常没有经过指令调度优化。可以尝试:

  • 在 prompt 里明确要求“使用 vectorized memory access”、“最小化寄存器使用”、“使用 shared memory 做数据复用”。
  • 生成后手动跑一遍 ptxas 的优化选项(-O3)。
  • 把性能瓶颈(比如“occupancy 只有 25%”)反馈给 LLM,让它调整寄存器声明。

5.4 常见错误速查表

错误现象可能原因排查方法
ptxas 报语法错误指令拼写、寄存器类型不匹配检查报错行,对照 PTX ISA 文档
编译通过但段错误缺少边界检查、地址计算错误用最小测试用例(n=1)逐步放大
结果部分正确线程索引计算错、内存空间限定符错打印每个线程的 tid 和访问地址
性能远低于预期未使用 vectorized load、寄存器过多看 ptxas 的 verbose 输出,检查寄存器数
同一 prompt 结果不同采样随机性降低 temperature,固定 seed

避坑技巧:在让 LLM 生成 PTX 之前,先让它生成一份“步骤说明”,你人工检查步骤对不对,再让它翻译成 PTX。这样能提前发现逻辑错误,省去反复编译测试的时间。

6. 这个方向的实际价值与边界

6.1 适合用 LLM 生成 PTX 的场景

不是所有场景都适合让 LLM 写 PTX。我总结了几类适合的:

  • 模板化 kernel 的快速原型:比如你要试 10 种不同的 reduction 变体,手写太慢,让 LLM 批量生成再筛选。
  • 教学和实验:想理解某个操作在 PTX 层面长什么样,让 LLM 生成一个参考实现。
  • 编译器覆盖不到的角落:某些非常规的 memory layout 或者同步模式,传统编译器优化不好,LLM 可能给出更直接的表达。

6.2 不适合的场景

  • 生产环境的性能关键路径:LLM 生成的 PTX 性能不稳定,不能保证每次都达到手写水平。
  • 需要严格正确性保证的场景:概率模型没有正确性保证,必须配合大量测试。
  • 超大规模 kernel:PTX 代码量大了之后,LLM 容易“忘记”前面的声明,生成不一致的代码。

6.3 和现有工具链的融合方式

我的看法是,短期内 LLM-as-compiler 不会替代 NVCC 或 Triton,而是作为一种补充工具存在。可能的融合方式:

  • 作为 IDE 插件:你写 CUDA 代码,它实时显示对应的 PTX,帮你理解编译器做了什么。
  • 作为优化助手:你有一个性能不好的 kernel,它生成几个 PTX 变体,你挑最快的。
  • 作为学习工具:你想学 PTX,它根据你的描述生成示例,你对照学习。

论文里提到的“AI lowering”这个概念,我觉得长期来看是有价值的。传统编译器的 lowering 是规则驱动的,规则是人写的,覆盖不了所有情况。LLM 的 lowering 是数据驱动的,理论上能覆盖更广的模式空间。但前提是,我们得有可靠的验证机制来兜底。

我个人在实际折腾这个方向的过程中,最大的体会是:LLM 生成 PTX 的瓶颈不在“生成”,而在“验证”。生成一段 PTX 只要几秒钟,但验证它是否正确、是否高效,可能需要几分钟甚至更久。所以整个流程的设计重心应该放在自动化验证上,而不是一味追求生成质量。把验证做扎实了,哪怕生成质量一般,也能通过迭代筛选出可用的结果。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询