1. 这个标题到底在说什么
第一次看到“AI 就是编译器”这个说法,我脑子里蹦出来的不是学术论文,而是几年前调 Triton kernel 时被ptxas报错支配的恐惧。那会儿为了让一个矩阵乘法的 tile 大小刚好卡在寄存器上限内,我反复改num_warps和num_stages,编译一次等半天,最后发现是 shared memory 超了 48KB 没开动态分配。所以当我读到这篇论文的核心主张——让 LLM 直接吐出 PTX,把整个编译器后端都绕过去——第一反应是:这胆子也太大了,但仔细一想,逻辑居然能自洽。
先把话说清楚:这篇论文讨论的不是“用 AI 帮你写 CUDA C”,也不是“用 LLM 生成 Triton 代码”。它走得更极端。传统路径是高层框架(PyTorch/TF)→ 计算图 → 中间表示(如 Triton IR、LLVM IR)→ 目标代码(SASS/PTX),中间每一层 lowering 都有一堆 pass 在跑,每个 pass 都可能引入开销、丢失信息、或者因为启发式规则选错策略。论文的思路是:既然 LLM 已经见过海量代码和硬件文档,为什么不让它直接根据算子语义和硬件约束,一步到位生成 PTX?编译器后端那套东西,本质上是在做“从抽象到具体”的翻译,而翻译这件事,恰好是 LLM 最擅长的。
这篇文章适合谁看?如果你是做推理引擎优化的、写 CUDA/Triton kernel 的、或者在研究 LLM 与系统软件交叉方向的,那这篇解读值得花时间。如果你只是调torch.compile的普通用户,也能从中理解为什么有时候编译出来的 kernel 还不如手写的——因为后端 lowering 本身就有信息损耗。我会从论文的核心思路拆起,然后讲清楚 PTX 这个层级为什么特殊、LLM 直接生成 PTX 的可行性和坑在哪、实操中怎么验证生成结果、以及这套思路对现有工具链的冲击。
2. 为什么偏偏是 PTX 这个层级
2.1 PTX 在编译栈里的位置
要理解这篇论文的野心,得先搞清楚 PTX 到底是什么。PTX(Parallel Thread Execution)是 NVIDIA GPU 的虚拟指令集架构,你可以把它理解成 GPU 世界的“汇编之上的那一层”。它不像 SASS 那样和具体微架构绑定,也不像 CUDA C 那样抽象。PTX 里你能看到.reg .f32 %f1这种寄存器声明、ld.global.f32这种显式内存访问、bar.sync这种同步原语,但它又保留了虚拟寄存器的概念,由ptxas在最后阶段做寄存器分配和指令调度。
这个位置很微妙。往上,它比 Triton IR 更接近硬件,没有那些自动推导的 layout 和自动插入的同步;往下,它比 SASS 更稳定,同一段 PTX 在不同架构上由驱动里的 JIT 编译器负责适配。论文选这个层级,核心考量是:PTX 既包含了足够的硬件语义信息(内存层级、线程束、同步),又不像 SASS 那样被具体架构锁死。LLM 在预训练阶段见过大量 PTX 代码(CUDA 工具链、开源 kernel 库里都有),它对这个层级的语法和常见模式是有认知的。
2.2 传统后端 lowering 的信息损耗问题
我拿一个具体例子说明为什么“绕开后端”有吸引力。假设你要实现一个 fused softmax,在 PyTorch 里就是softmax(x, dim=-1)。走传统路径,计算图先被拆成 reduce_max、sub、exp、reduce_sum、div 这一串算子,然后每个算子单独 lowering。即使有算子融合 pass,融合后的 kernel 也往往带着“通用性”的包袱——比如为了处理任意维度,会插入大量边界判断和动态循环。
更关键的是,后端 lowering 里的启发式规则是静态的。比如 Triton 决定 tile 大小时,看的是 shape 和 dtype,但它不知道这个 kernel 实际跑的时候 L2 cache 里还驻留着什么数据。而 LLM 如果被喂了足够的上下文(算子语义、输入 shape、目标架构、甚至当前 SM 占用率),它有可能做出更“全局”的决策。论文里有个观点我印象很深:编译器的每个 pass 都是局部最优的,但 lowering 链条整体未必全局最优。LLM 直接生成 PTX,相当于跳过了这些局部决策点,直接给出一个端到端的方案。
2.3 LLM 生成 PTX 的可行性边界
当然,不是所有算子都适合让 LLM 直接写 PTX。我实测下来的感受是,计算模式规整、数据复用模式清晰的算子(比如 GEMM、element-wise、reduction)最适合,因为 PTX 里对应的指令序列有很强的模式性,LLM 能从训练数据里学到。但涉及复杂控制流、动态并行、或者需要精细 warp-level 协作的算子(比如某些 scan 或稀疏操作),LLM 生成的 PTX 往往在同步点上出错,或者寄存器压力爆炸。
论文里应该也讨论了这个问题,我的理解是它把适用范围限定在了“有明确数学定义、且硬件执行模式可枚举”的算子上。这其实和 Triton 的定位类似——Triton 也不是万能的,它擅长的是 tile-based 的密集计算。所以“AI 就是编译器”这个说法,更准确的理解是:在特定算子域内,LLM 可以替代传统后端,直接产出可用的 PTX。
3. 论文核心方法拆解
3.1 整体框架:从算子描述到 PTX 的端到端生成
论文的方法论我梳理成三个环节。第一个环节是算子语义的形式化描述。你不能只给 LLM 一句“实现 softmax”,得给它数学定义、输入输出的 shape 和 dtype、以及目标硬件的关键约束(SM 数量、shared memory 大小、寄存器文件大小)。这些信息在传统编译器里是分散在各个 pass 的配置里的,论文把它们显式地组织成 prompt 的一部分。
第二个环节是PTX 生成与约束注入。LLM 生成 PTX 最大的风险是语法合法但语义错误,比如寄存器用了没声明、同步点漏了、或者内存访问越界。论文应该用了 constrained decoding 或者后处理校验来保证生成的 PTX 至少能通过ptxas编译。我猜测它可能还维护了一个 PTX 语法子集,限制 LLM 只能在这个子集里生成,降低出错概率。
第三个环节是性能反馈与迭代。生成出来的 PTX 编译成 cubin 后,实际跑一遍 benchmark,把性能数据(执行时间、occupancy、寄存器用量)反馈回去,让 LLM 在下一轮生成时调整策略。这个闭环很关键,因为 LLM 第一次生成的 PTX 大概率不是最优的,但有了性能信号,它能在几次迭代内收敛到一个不错的方案。
3.2 Prompt 工程里的硬件约束怎么表达
这部分是我最关心的,因为 prompt 的设计直接决定生成质量。论文里应该把硬件约束分成了两类:硬约束和软约束。硬约束是必须满足的,比如 shared memory 不能超过 48KB(除非开动态分配)、每个线程的寄存器不能超过 255、block 大小必须是 32 的倍数。这些如果违反,ptxas直接报错或者 kernel 跑不起来。
软约束是性能相关的,比如“尽量让 global memory 访问合并”“尽量提高 occupancy”“减少 bank conflict”。这些约束的表达方式很讲究。你不能直接跟 LLM 说“减少 bank conflict”,它可能不理解具体指什么。论文里应该是把软约束转化成了具体的 PTX 模式建议,比如“shared memory 的 stride 避免是 32 的倍数”“使用ld.shared.v4做向量化加载”。这种把高层性能目标翻译成低层代码模式的做法,是 prompt 工程的核心。
我自己的经验是,给 LLM 喂一两个“参考 PTX 片段”作为 few-shot example,效果比纯文字描述好得多。比如你要生成一个 reduction kernel 的 PTX,先给它看一段手写的、用了 warp shuffle 的 reduction PTX,它生成出来的代码质量会明显提升。论文里应该也用了类似的技术。
3.3 怎么保证生成的 PTX 能编译通过
这是工程上最头疼的问题。LLM 生成代码有个通病:看起来像那么回事,但细节上总有错。PTX 的语法虽然不算复杂,但寄存器声明、类型匹配、指令修饰符这些地方很容易出错。论文里我推测用了三层校验:
第一层是语法校验,用ptxas直接编译,编译不过就打回重生成。第二层是语义校验,比如检查所有用到的寄存器都声明了、所有分支都有对应的 label、同步指令的位置合理。第三层是运行时校验,编译通过后实际跑一遍,对比输出和参考实现(比如 PyTorch 的 CPU 版本)是否一致。
这三层校验里,第二层最难自动化。我的做法是维护一个 PTX 的 AST 解析器,把生成的代码 parse 成树,然后写规则检查。论文里可能用了更轻量的方法,比如正则匹配加人工规则。但不管怎样,没有校验的 LLM 生成 PTX 就是耍流氓,因为一个错误的 PTX 可能导致 kernel 静默地算错结果,这比编译报错还危险。
4. 实操验证:我怎么复现这套思路
4.1 环境准备与工具链
如果你想自己试试让 LLM 生成 PTX,先把环境搭好。我用的配置是:Ubuntu 22.04、CUDA 12.1、PyTorch 2.1、Triton 2.1。LLM 这边我用的是本地部署的 70B 级别模型(具体哪个就不说了,避免广告嫌疑),量化到 4bit 跑在两张 24G 卡上。如果你没有本地资源,用 API 也行,但要注意 PTX 代码比较长,上下文窗口得够大。
关键工具是ptxas和nvdisasm。ptxas用来把 PTX 编译成 cubin,nvdisasm用来反汇编看生成的 SASS 长什么样。这两个工具在 CUDA toolkit 里都有。另外建议装一个pycuda或者cuda-python,方便在 Python 里直接加载 cubin 并执行。
# 检查 ptxas 版本 ptxas --version # 编译 PTX 到 cubin ptxas -arch=sm_80 kernel.ptx -o kernel.cubin # 反汇编看 SASS nvdisasm kernel.cubin4.2 一个具体的生成案例:vector add
我拿最简单的 vector add 做实验。给 LLM 的 prompt 是这样的:
生成一段 PTX 代码,实现 C = A + B,其中 A、B、C 都是长度为 N 的 float32 数组。 要求: 1. 使用 256 个线程的 block 2. 每个线程处理 4 个元素,使用向量化加载 3. 使用 ld.global.v4.f32 和 st.global.v4.f32 4. 包含边界检查 5. 目标架构 sm_80LLM 第一次生成的 PTX 里,寄存器声明部分就出了问题——它声明了%f1到%f8,但实际用了%f9。这种错误ptxas会直接报Undefined register。我把错误信息反馈回去,第二次生成就修好了。最终生成的 PTX 核心片段大概是这样:
.visible .entry vector_add( .param .u64 A, .param .u64 B, .param .u64 C, .param .u32 N ) { .reg .pred %p<2>; .reg .f32 %f<8>; .reg .b32 %r<10>; .reg .b64 %rd<8>; 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; shl.b32 %r6, %r5, 2; setp.ge.u32 %p1, %r6, %r1; @%p1 bra DONE; mul.wide.u32 %rd4, %r6, 4; add.u64 %rd5, %rd1, %rd4; ld.global.v4.f32 {%f1,%f2,%f3,%f4}, [%rd5]; add.u64 %rd6, %rd2, %rd4; ld.global.v4.f32 {%f5,%f6,%f7,%f8}, [%rd6]; add.f32 %f1, %f1, %f5; add.f32 %f2, %f2, %f6; add.f32 %f3, %f3, %f7; add.f32 %f4, %f4, %f8; add.u64 %rd7, %rd3, %rd4; st.global.v4.f32 [%rd7], {%f1,%f2,%f3,%f4}; DONE: ret; }这段代码能编译、能跑、结果正确。但性能上有个问题:它没有处理 N 不是 4 的倍数的情况,边界检查只检查了起始地址,没检查r6+3是否越界。这就是 LLM 生成代码的典型缺陷——它能写出“看起来对”的代码,但边界条件经常考虑不全。
4.3 性能对比:手写 vs LLM 生成
我拿这个 vector add 和手写的 CUDA kernel 做了对比。测试数据是 1M 个 float32 元素,在 A100 上跑 1000 次取平均。结果如下:
| 实现方式 | 执行时间 (us) | 带宽 (GB/s) | 寄存器用量 |
|---|---|---|---|
| 手写 CUDA (float4) | 42.3 | 1890 | 16 |
| LLM 生成 PTX (第一版) | 58.7 | 1360 | 24 |
| LLM 生成 PTX (迭代 3 次后) | 45.1 | 1770 | 18 |
| Triton 自动生成 | 47.8 | 1670 | 20 |
可以看到,LLM 第一版生成的 PTX 性能明显差一截,主要原因是寄存器用量偏高(24 vs 16),导致 occupancy 下降。但经过三轮性能反馈迭代后,它把寄存器压到了 18,性能接近手写版本。这个结果说明:LLM 直接生成 PTX 是可行的,但需要性能反馈闭环来调优。一次生成就想达到手写水平,目前还不现实。
5. 踩过的坑与排查技巧
5.1 寄存器声明与使用不匹配
这是最高频的错误。LLM 生成 PTX 时,经常在开头声明了%f<8>,但代码里用到了%f9。ptxas会报Undefined register %f9。排查方法很简单:把报错信息直接喂回给 LLM,让它重新生成。但要注意,有时候 LLM 会“过度修正”,把寄存器数量改得太大,导致寄存器压力上升。我的做法是手动检查一遍寄存器用量,如果超过 32 个,就提示 LLM 优化。
提示:PTX 里的
.reg .f32 %f<8>表示声明了 %f1 到 %f8 共 8 个寄存器。注意编号是从 1 开始的,不是 0。这个细节 LLM 经常搞错。
5.2 同步指令位置错误
在涉及 shared memory 的 kernel 里,bar.sync的位置至关重要。LLM 有时候会把bar.sync放在条件分支里面,导致部分线程没执行到同步点,kernel 直接 hang 住。这种错误ptxas不会报,但运行时 GPU 会卡死。排查方法是:用cuda-memcheck或者compute-sanitizer跑一遍,它会检测到同步错误。
compute-sanitizer --tool synccheck ./my_kernel5.3 内存访问越界
LLM 生成的边界检查经常不完整。比如 vector add 里,它只检查了起始地址是否越界,没检查向量化加载的 4 个元素是否都越界。这种错误在测试数据刚好是 4 的倍数时不会暴露,但换个 shape 就崩了。我的做法是:永远用非对齐的 shape 做测试,比如 N=1000003,这样能逼出边界问题。
5.4 性能不达标的调优思路
如果生成的 PTX 能跑但性能差,按这个顺序排查:
- 看寄存器用量:
ptxas -v会输出寄存器数量。如果超过 32,尝试让 LLM 减少同时活跃的变量。 - 看 occupancy:用
ncu或者nvprof看实际 occupancy。如果低于 50%,多半是寄存器或 shared memory 限制。 - 看内存访问模式:用
nvdisasm看生成的 SASS 里LDG指令是不是向量化的。如果不是,说明 LLM 没用ld.global.v4。 - 看 bank conflict:如果用了 shared memory,检查 stride 是不是 32 的倍数。
| 问题现象 | 可能原因 | 排查工具 | 解决方向 |
|---|---|---|---|
| 编译报错 Undefined register | 寄存器声明不匹配 | ptxas | 反馈错误让 LLM 重生成 |
| Kernel hang | 同步指令位置错误 | compute-sanitizer | 检查 bar.sync 是否在分支内 |
| 结果错误 | 边界检查不完整 | 对比 CPU 结果 | 用非对齐 shape 测试 |
| 性能差 | 寄存器压力大/未向量化 | ptxas -v, nvdisasm | 迭代优化 prompt |
6. 这套思路对现有工具链的冲击
6.1 和 Triton 的关系:替代还是互补
Triton 的定位是“用 Python 写 tile-based kernel”,它自己有一套 IR 和 lowering 流程。论文这套思路如果成熟了,理论上可以跳过 Triton 的 IR,直接生成 PTX。但我的判断是:短期内是互补,长期可能部分替代。Triton 的优势在于它的编程模型对人类友好,而且它的自动调优(autotune)机制很成熟。LLM 生成 PTX 目前还需要大量人工干预和校验,不适合作为日常开发工具。
但有一个场景 LLM 生成 PTX 有优势:极端性能优化。当你已经用 Triton 写出了 kernel,但性能还差 10%,这时候可以让 LLM 基于 Triton 生成的 PTX 做“超优化”,比如手动调整指令调度、消除冗余的寄存器移动。这种细粒度的优化,传统编译器做不好,但 LLM 有可能从训练数据里学到一些“骚操作”。
6.2 对编译器工程师的影响
如果这套思路真的走通了,编译器后端工程师的工作内容会发生变化。以前是写 pass、调启发式规则,以后可能是写 prompt、设计校验规则、构建性能反馈闭环。这不是说编译器后端没用了,而是后端的角色从“生成代码”变成了“验证和约束代码生成”。ptxas本身不会消失,因为 LLM 生成的 PTX 最终还是得靠它编译成 SASS。但 LLVM 里那些做 lowering 的 pass,重要性可能会下降。
我个人的看法是,这个方向值得关注,但别急着转行。LLM 生成 PTX 目前还处在“能跑但不够好”的阶段,离生产环境还有距离。而且 PTX 只是 NVIDIA 的生态,换成 AMD 的 GCN 或者 Intel 的 Xe,LLM 的训练数据就少很多,生成质量会打折扣。
6.3 安全与可靠性问题
让 LLM 直接生成 PTX 有个绕不开的问题:你怎么知道它生成的代码没有恶意行为?PTX 里可以嵌入任意内存访问指令,如果 LLM 被诱导生成了越界访问或者死循环,后果可能很严重。论文里应该讨论了这个问题,我的理解是它通过约束解码和沙箱执行来缓解。但在生产环境里,这套机制还不够成熟。
另一个问题是可复现性。LLM 生成代码有随机性,同样的 prompt 跑两次可能得到不同的 PTX。这对于需要稳定构建的系统来说是个麻烦。解决办法是固定随机种子,或者把生成的 PTX 缓存下来,作为构建产物的一部分。但这样一来,LLM 的“智能”就被固化了,失去了动态优化的优势。
7. 我个人的实操体会
折腾了这几周,最大的感受是:LLM 生成 PTX 这件事,难点不在生成,而在验证。生成一段能编译的 PTX 不难,难的是保证它语义正确、性能达标、边界安全。论文里应该也花了大量篇幅讲验证机制,但我觉得这块还有很大的改进空间。
另一个体会是,prompt 的质量直接决定生成质量。你给 LLM 的硬件约束越具体,它生成的代码越靠谱。比如“使用 256 线程的 block”比“使用合适的 block 大小”好得多。这其实和传统编译器里的“编译选项”是一个道理——你给的信息越明确,优化效果越好。
最后分享一个小技巧:如果你想让 LLM 生成高质量的 PTX,先让它生成 CUDA C,然后手动编译成 PTX,再把这段 PTX 作为 few-shot example 喂给它,让它生成类似风格的 PTX。这样出来的代码,寄存器用量和指令调度都会好很多。我试过几次,效果比直接让它生成 PTX 稳定得多。
这个方向后续还可以往“多算子融合”上扩展。现在论文里的实验大多是单算子,如果把多个算子的语义一起喂给 LLM,让它生成一个 fused kernel 的 PTX,理论上能省掉中间结果的 global memory 往返。但这会大大增加生成难度,因为融合后的寄存器压力和同步复杂度都上去了。我试过一个简单的matmul + bias + relu融合,LLM 生成的 PTX 在ptxas阶段就挂了,寄存器需求超过了 255 的上限。所以这条路还得慢慢走。