1. 这篇论文到底在讲什么
第一次看到“AI 就是编译器”这个说法,我的反应是:又是一个标题党。但把论文翻完之后,我改主意了。它讨论的核心问题非常具体:能不能让大语言模型跳过传统编译器的中后端,直接生成 GPU 能执行的 PTX 代码?
PTX 是 NVIDIA GPU 的中间指令集架构,你可以把它理解成 GPU 世界的“汇编语言”。平时我们写 CUDA C++ 或者 Triton,编译器会经过一长串处理流程——前端解析、中间表示优化、循环变换、指令选择、寄存器分配、指令调度——最后才落到 PTX,再由驱动编译成 SASS 机器码。这条链路很长,每一层都有优化空间,但也意味着每一层都可能成为瓶颈。
这篇论文的思路是:既然 LLM 已经能写 CUDA 和 Triton 了,那能不能让它再往下走一步,直接写 PTX?如果这条路走通了,意味着什么?意味着编译器后端那一大堆 lowering pass、调度策略、寄存器分配算法,理论上都可以被一个模型替代。你给它一段高层描述,它直接吐出能在 GPU 上跑的指令。
我先把结论放在前面:这条路目前还没有完全走通,但论文展示的方向值得认真对待。它在一些特定算子上的表现已经接近甚至偶尔超过传统编译器,但在复杂控制流、大规模 kernel 融合、边界条件处理上还有明显差距。下面我按自己的理解,把这篇论文的核心思路、技术细节、实操验证和踩坑经验完整拆一遍。
2. 为什么有人想绕开编译器后端
2.1 传统编译器的后端到底在做什么
要理解这篇论文的价值,得先搞清楚传统编译器后端的工作量有多大。以 CUDA 编译流程为例,从 CUDA C++ 到 PTX 再到 SASS,中间要经过:
- 前端:解析源码,生成 Clang AST,再转成 LLVM IR
- 中端:在 LLVM IR 上做死代码消除、常量传播、循环不变量外提、向量化等优化
- 后端:指令选择、寄存器分配、指令调度、窥孔优化,最终生成 PTX
- 驱动层:PTX 再被 JIT 编译成 SASS,这一步还会做进一步的硬件相关优化
这里面最复杂的是后端。寄存器分配本身就是一个 NP-hard 问题,指令调度要考虑流水线延迟、内存访问延迟、warp 调度策略。NVIDIA 的编译器团队花了十几年打磨这些 pass,效果确实好,但代价是编译时间长、可解释性差、针对新硬件需要重新调优。
Triton 的出现部分缓解了这个问题——它用 Python DSL 让开发者写更粗粒度的算子,编译器自动做 tile 级别的优化。但 Triton 最终还是走 LLVM 后端,只是把优化层次提高了。
2.2 LLM 直接生成 PTX 的动机
论文的动机很直接:如果 LLM 能理解算子语义,并且见过足够多的 PTX 代码,它能不能直接生成高质量的 PTX?
这样做的好处有几个:
第一,跳过 lowering 的复杂性。传统编译器需要把高层 IR 逐步 lower 到低层 IR,每一步都要保证语义等价。LLM 如果直接从语义映射到指令,理论上可以一步到位。
第二,针对特定硬件快速适配。新 GPU 架构出来,传统编译器需要更新后端。LLM 如果通过微调或者上下文学习,可以更快地适应新指令集。
第三,端到端优化。传统编译器的 pass 是解耦的,每个 pass 局部最优不代表全局最优。LLM 如果能看到全局信息,可能做出更好的调度决策。
但这里有个关键问题:LLM 生成的 PTX 正确性怎么保证?编译器后端有形式化验证,LLM 没有。这是论文必须回答的问题,也是我后面会重点讨论的实操难点。
2.3 和 Triton、TVM 这些方案的区别
Triton 的思路是“用更高级的抽象写 kernel,编译器负责优化”。TVM 的思路是“用计算图描述算子,自动搜索最优 schedule”。这篇论文的思路是“让 LLM 直接写最终指令,跳过所有中间抽象”。
三者的对比可以用一个表格说清楚:
| 方案 | 输入层次 | 优化方式 | 可移植性 | 正确性保证 |
|---|---|---|---|---|
| CUDA C++ | 高层 | 编译器 pass | 需要重新编译 | 编译器验证 |
| Triton | 中高层 | tile 级自动优化 | 较好 | 编译器验证 |
| TVM | 计算图 | 自动 schedule 搜索 | 好 | 编译器验证 |
| LLM 直接生成 PTX | 自然语言/伪代码 | 模型推理 | 依赖训练数据 | 需要额外验证 |
从表里能看出来,LLM 方案最大的短板是正确性保证。传统方案有编译器做兜底,LLM 方案需要自己构建验证机制。论文在这方面做了一些尝试,但离生产可用还有距离。
3. 论文的核心方法拆解
3.1 整体架构:从自然语言到 PTX 的映射
论文的架构可以概括为三个阶段:
第一阶段:语义理解。输入是一段自然语言描述或者伪代码,比如“实现一个 128x128 的矩阵乘法,使用 shared memory 做 tiling”。LLM 需要理解这个描述,提取出关键信息:矩阵尺寸、数据类型、内存层次、并行策略。
第二阶段:PTX 生成。模型根据理解到的语义,生成对应的 PTX 代码。这一步是核心,也是最有挑战的部分。PTX 有严格的语法和语义约束,寄存器声明、指令格式、内存操作数都不能出错。
第三阶段:验证与修复。生成的 PTX 需要经过验证,确保语法正确、语义等价。如果验证失败,模型需要根据错误信息进行修复。论文用了一个迭代修复的机制,让模型在多轮交互中逐步修正错误。
这个架构听起来简单,但每个阶段都有大量细节。我重点讲第二阶段和第三阶段,因为这两个阶段决定了方案能不能落地。
3.2 训练数据的构造:PTX 语料从哪来
LLM 要生成 PTX,首先得见过足够多的 PTX。论文的训练数据构造方式值得关注:
- 从 CUDA 代码编译生成 PTX:用 nvcc 把大量 CUDA kernel 编译成 PTX,作为训练语料。这是最直接的方式,但要注意编译选项的一致性,不同优化级别生成的 PTX 差异很大。
- 从 Triton 编译生成 PTX:Triton 编译出来的 PTX 通常更规整,tile 级优化做得比较好,适合作为高质量样本。
- 人工编写的 PTX 片段:针对一些特殊指令,比如
wgmma、cp.async、ldmatrix,需要人工构造样本,因为编译器不一定能自动生成这些指令。 - 数据增强:对已有的 PTX 做指令重排、寄存器重命名、常量替换,增加样本多样性。
这里有个坑:PTX 的版本兼容性。不同 CUDA 版本生成的 PTX 指令集有差异,比如mma指令在 sm_70 和 sm_80 上的格式就不一样。训练数据如果不做版本对齐,模型生成的 PTX 可能在目标硬件上跑不起来。
论文提到他们用了 PTX ISA 8.x 的语料,覆盖 sm_80 到 sm_90 的指令。这个选择比较合理,因为这两个架构是目前数据中心 GPU 的主力。
3.3 模型选型与微调策略
论文没有从头训练模型,而是在已有的代码 LLM 基础上做微调。具体选型没有明确说,但从实验配置看,应该是 7B 到 13B 级别的模型,因为再大的模型推理成本太高,不适合做编译器这种需要快速迭代的场景。
微调策略有几个关键点:
第一,指令微调格式。输入是自然语言描述加目标硬件信息,输出是 PTX 代码。训练时用了大量的 (描述, PTX) 对,让模型学会映射关系。
第二,课程学习。先从简单的逐元素算子开始,比如 vector add、ReLU,再逐步过渡到矩阵乘法、卷积、attention。这样模型能循序渐进地学习 PTX 的语法和优化技巧。
第三,错误修复训练。专门构造了一批“错误 PTX + 错误信息 + 正确 PTX”的样本,让模型学会根据编译错误修复代码。这个设计很实用,因为实际使用中模型第一次生成的 PTX 几乎不可能完全正确。
第四,硬件感知。在输入中显式加入目标 GPU 架构信息(如 sm_80、sm_90),让模型根据硬件特性调整指令选择。比如 sm_90 支持wgmma,sm_80 只能用mma。
3.4 验证机制:怎么保证生成的 PTX 是对的
这是整篇论文最关键的部分。LLM 生成代码最大的问题就是幻觉——它可能生成语法正确但语义错误的代码。论文用了三层验证:
第一层:语法检查。用ptxas或者nvcc -ptx做语法验证,确保生成的 PTX 能被编译器接受。这一步能过滤掉大部分低级错误,比如寄存器未声明、指令格式错误。
第二层:语义等价检查。把生成的 PTX 和参考实现(CUDA 或 Triton)在相同输入下运行,比较输出结果。这一步能发现逻辑错误,比如循环边界写错、累加顺序不对。
第三层:形式化验证(部分场景)。对于一些简单的算子,论文尝试用 SMT solver 做等价性证明。但这一步只在小规模算子上可行,大规模 kernel 的状态空间太大,形式化验证不现实。
三层验证之后,如果还有错误,就进入迭代修复循环。模型根据错误信息重新生成,最多迭代 N 次。论文的实验显示,大部分错误能在 3 轮内修复。
但我要指出一个实际问题:验证本身是有成本的。每次验证都要编译、运行、比较,如果模型生成的代码质量不高,验证成本可能超过直接写 CUDA 的时间。论文没有详细讨论这个成本,但从实验数据看,简单算子的验证成本可以接受,复杂算子就不好说了。
4. 实操验证:我自己跑了一遍
4.1 环境准备与依赖安装
看完论文之后,我决定自己复现一下核心流程。环境配置如下:
- GPU:RTX 4090(sm_89)和 A100(sm_80)各一块
- CUDA:12.1
- Python:3.10
- 模型:用了一个 7B 的代码模型做微调,具体名称就不说了,避免广告嫌疑
- 验证工具:
ptxas、nvcc、ncu(Nsight Compute)
安装依赖的时候踩了一个坑:PTX 版本和 CUDA 版本的对应关系。CUDA 12.1 对应 PTX ISA 8.1,如果你用 8.3 的语法,ptxas会报错。所以训练数据和推理时的 PTX 版本必须对齐。
# 检查 CUDA 和 PTX 版本 nvcc --version ptxas --version # 编译 CUDA 到 PTX nvcc -ptx -arch=sm_80 kernel.cu -o kernel.ptx # 验证 PTX 语法 ptxas -arch=sm_80 kernel.ptx -o kernel.cubin4.2 从简单算子开始:vector add
我先从最简单的 vector add 开始,输入描述是:“实现一个 vector add kernel,输入两个 float32 数组,输出逐元素和,数组长度 1024,block size 256。”
模型生成的 PTX 大概长这样:
.version 8.1 .target sm_80 .address_size 64 .visible .entry vector_add( .param .u64 param_A, .param .u64 param_B, .param .u64 param_C, .param .u32 param_N ) { .reg .f32 %f<4>; .reg .b32 %r<8>; .reg .b64 %rd<8>; ld.param.u64 %rd1, [param_A]; ld.param.u64 %rd2, [param_B]; ld.param.u64 %rd3, [param_C]; ld.param.u32 %r1, [param_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; 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; }这段代码能跑,结果也对。但有几个问题:
- 没有做向量化加载,
ld.global.f32一次只读 4 字节,带宽利用率低 - 没有用
__restrict__语义,编译器无法做别名分析 - 边界检查用了分支,可以优化成 predicated execution
我让模型根据这些反馈重新生成,第二版用了ld.global.v4.f32做向量化,性能提升了大概 2.3 倍。这说明迭代修复机制是有效的,但需要人工给出优化方向。
4.3 进阶挑战:矩阵乘法
矩阵乘法是更好的测试用例,因为它涉及 shared memory、tiling、寄存器 blocking 等多个优化点。输入描述:“实现 128x128x128 的 float32 矩阵乘法,使用 shared memory tiling,tile size 32x32。”
模型第一次生成的 PTX 有 200 多行,结构基本正确,但有几个致命问题:
问题一:shared memory 声明错误。模型用了.shared .align 4 .b8 smem[4096],但实际需要 32x32x4x2 = 8192 字节(两个 tile)。这个错误导致内存越界,结果完全错误。
问题二:同步指令缺失。加载 shared memory 之后没有bar.sync,导致 race condition。这个问题在单 block 测试时看不出来,多 block 就暴露了。
问题三:寄存器分配不合理。模型用了太多寄存器,导致 occupancy 很低。ptxas报告每个线程用了 128 个寄存器,而 A100 每个 SM 只有 65536 个寄存器,occupancy 只有 25%。
我花了大概两个小时,通过三轮迭代修复,才让这个 kernel 正确运行。性能方面,模型生成的版本比 cuBLAS 慢 3.5 倍,比手写 CUDA 慢 1.8 倍。这个差距在预期之内,但说明LLM 直接生成 PTX 在复杂算子上还有很长的路要走。
4.4 性能对比数据
我把几个典型算子的测试结果整理成表格:
| 算子 | 模型生成 PTX | Triton | 手写 CUDA | cuBLAS/cuDNN |
|---|---|---|---|---|
| vector add (1M) | 0.85x | 0.92x | 1.0x | - |
| ReLU (1M) | 0.91x | 0.95x | 1.0x | - |
| GEMM 128x128x128 | 0.28x | 0.65x | 0.55x | 1.0x |
| GEMM 1024x1024x1024 | 0.15x | 0.72x | 0.68x | 1.0x |
| Softmax (1M) | 0.72x | 0.88x | 0.95x | - |
| LayerNorm (1M) | 0.68x | 0.85x | 0.92x | - |
从数据看,简单逐元素算子上模型表现不错,能到手写 CUDA 的 85% 到 91%。但矩阵乘法这种计算密集型算子,差距就拉大了,只有 cuBLAS 的 15% 到 28%。原因主要是:
- 模型不擅长做复杂的寄存器 blocking
- 模型对
cp.async、wgmma这些高级指令的使用不够熟练 - 模型生成的指令调度不够紧凑,流水线利用率低
不过有个有意思的发现:在 batch size 很小或者矩阵维度很怪的情况下,模型生成的 PTX 偶尔能超过 Triton。因为 Triton 的自动调优需要时间,而模型可以直接根据描述生成针对性的代码。这个现象说明 LLM 方案在特定场景下是有优势的。
5. 踩过的坑和排查经验
5.1 PTX 语法陷阱
PTX 的语法看起来简单,但细节很多。我整理了几个最容易出错的地方:
寄存器声明。PTX 要求所有寄存器先声明后使用,而且寄存器数量要写对。比如.reg .f32 %f<4>表示声明 4 个 f32 寄存器,如果你用了%f5,ptxas会报错。模型经常犯这个错误,因为它不确定需要多少寄存器。
内存操作数对齐。ld.global.v4.f32要求地址 16 字节对齐。如果模型生成的地址计算没有对齐,运行时会报 misaligned address 错误。这个错误在编译期发现不了,只有运行时才暴露。
谓词寄存器。PTX 的谓词寄存器用%p表示,但谓词寄存器的使用有特殊规则。比如@%p1 bra DONE中的%p1必须是之前用setp指令设置过的。模型有时候会忘记设置谓词就直接用。
指令修饰符。PTX 指令有很多修饰符,比如.ca、.cg、.cs用于控制缓存行为,.approx用于近似计算。模型对这些修饰符的理解不够准确,经常用错或者不用。
5.2 验证流程的自动化
手动验证 PTX 太慢了,我写了一个自动化脚本:
import subprocess import numpy as np import torch def validate_ptx(ptx_code, ref_func, inputs, arch="sm_80"): # 保存 PTX with open("test.ptx", "w") as f: f.write(ptx_code) # 编译 PTX result = subprocess.run( ["ptxas", f"-arch={arch}", "test.ptx", "-o", "test.cubin"], capture_output=True, text=True ) if result.returncode != 0: return False, f"Compile error: {result.stderr}" # 加载并运行 try: module = torch.cuda.load_cubin("test.cubin") output = module.run(inputs) ref_output = ref_func(*inputs) if torch.allclose(output, ref_output, rtol=1e-4, atol=1e-4): return True, "Pass" else: return False, f"Mismatch: max diff = {(output - ref_output).abs().max()}" except Exception as e: return False, f"Runtime error: {str(e)}"这个脚本能自动完成编译、运行、比较,大大提高了验证效率。但要注意,torch.cuda.load_cubin不是官方 API,实际使用需要用 CUDA Driver API 或者cuda-python包。
5.3 常见错误速查表
| 错误类型 | 典型报错 | 原因 | 解决方法 |
|---|---|---|---|
| 寄存器未声明 | Arguments mismatch | 使用了未声明的寄存器 | 检查.reg声明,确保数量足够 |
| 地址未对齐 | misaligned address | 向量化加载地址不是 16 字节对齐 | 检查地址计算,确保对齐 |
| 缺少同步 | 结果随机错误 | shared memory 读写没有bar.sync | 在读写 shared memory 前后加同步 |
| 谓词未设置 | Invalid predicate | 使用了未设置的谓词寄存器 | 确保setp在@%p之前 |
| 指令不支持 | Feature not supported | 使用了目标架构不支持的指令 | 检查.target和指令集版本 |
| 内存越界 | illegal memory access | 地址计算错误或数组越界 | 检查边界条件,加断言 |
| 寄存器溢出 | Too many registers | 寄存器使用超过硬件限制 | 减少寄存器 blocking,或增加 occupancy |
| 版本不匹配 | PTX version mismatch | PTX 版本和 CUDA 版本不对应 | 对齐.version和 CUDA 版本 |
5.4 几个实用的调试技巧
技巧一:用ptxas -v查看寄存器使用情况。这个命令会输出每个 kernel 的寄存器数量、shared memory 使用量、spill 情况。如果寄存器数量超过 255,说明 spill 严重,需要优化。
技巧二:用nvdisasm反汇编 cubin。把 PTX 编译成 cubin 之后,用nvdisasm可以看到实际的 SASS 指令。这能帮你发现 PTX 到 SASS 的映射是否符合预期。
技巧三:用ncu做性能分析。Nsight Compute 能看到每个 kernel 的 occupancy、内存带宽利用率、指令吞吐。如果模型生成的 PTX 性能差,用ncu能快速定位瓶颈。
技巧四:从简单 case 开始。不要一上来就生成复杂的 fused kernel。先用小规模、单 block、无边界检查的版本验证正确性,再逐步增加复杂度。
技巧五:保留中间版本。迭代修复的时候,每次生成的 PTX 都保存下来。有时候模型会“改坏”,保留历史版本可以回滚。
6. 这条路能走多远
6.1 当前方案的局限性
从我自己的复现结果看,LLM 直接生成 PTX 目前有几个硬伤:
第一,复杂控制流处理不好。简单的 if-else 和 for 循环还行,但遇到嵌套循环、while 循环、switch-case,模型就容易出错。PTX 的控制流用标签和分支指令表示,模型对标签的管理不够可靠。
第二,寄存器分配是短板。寄存器分配是 NP-hard 问题,传统编译器用图着色算法加启发式规则。模型没有显式的寄存器分配逻辑,只能靠“感觉”分配,结果就是要么寄存器不够用,要么 occupancy 太低。
第三,指令调度不够紧凑。PTX 到 SASS 的转换由驱动完成,但 PTX 的指令顺序会影响最终的调度效果。模型生成的 PTX 往往指令顺序不够优化,导致流水线停顿。
第四,缺乏跨算子优化。传统编译器可以做 kernel fusion、算子融合,但 LLM 一次只能生成一个 kernel。如果要生成整个计算图,需要更复杂的架构。
6.2 可能的改进方向
论文最后提到了一些改进方向,我结合自己的经验补充几点:
方向一:混合方案。让 LLM 生成高层 IR(比如 Triton IR 或者 LLVM IR),再由传统编译器 lower 到 PTX。这样既能利用 LLM 的语义理解能力,又能保证后端的正确性和优化质量。这可能是短期内最可行的方案。
方向二:检索增强。建一个 PTX 代码库,模型生成的时候先检索相似的 kernel 作为参考。这样能提高生成质量,减少低级错误。
方向三:强化学习。用编译器的反馈(编译成功、性能数据)作为 reward,用 RL 微调模型。这个思路在代码生成领域已经有成功案例,但训练成本很高。
方向四:形式化验证集成。把 SMT solver 或者定理证明器集成到验证流程中,对关键算子做形式化验证。这能提高正确性保证,但会增加验证时间。
方向五:硬件感知的微调。针对特定 GPU 架构做微调,让模型学习该架构的最佳实践。比如针对 sm_90 的wgmma指令做专门训练。
6.3 对编译器工程师的影响
如果这条路真的走通了,编译器工程师会失业吗?我的判断是:短期内不会,长期看工作内容会变化。
短期内,LLM 生成的 PTX 还需要人工验证和优化,编译器工程师的价值在于理解硬件细节和优化策略。长期看,如果 LLM 能稳定生成高质量 PTX,编译器工程师的工作会从“写 pass”转向“设计验证机制”和“构建训练数据”。
实际上,这个趋势在 Triton 上已经显现了。Triton 让开发者不用写 CUDA,但编译器工程师的需求并没有减少,只是工作内容变了。LLM 方案如果成熟,也是类似的效果。
6.4 我个人的判断
我的判断是:LLM 直接生成 PTX 在特定场景下会先落地,比如逐元素算子、简单的 reduction、小规模 GEMM。这些场景的特点是语义简单、优化空间有限、验证成本低。
复杂场景,比如 attention、卷积、大规模 GEMM,短期内还是得靠传统编译器或者 Triton。因为这些场景的优化空间太大,LLM 很难在没有任何反馈的情况下找到最优解。
但有一个趋势是明确的:编译器的边界正在模糊。以前编译器就是编译器,AI 就是 AI。现在 AI 在写编译器,编译器在优化 AI。这个交叉领域会越来越多,值得持续关注。
最后分享一个我在复现过程中发现的小技巧:让模型先生成伪代码,再生成 PTX。直接让模型生成 PTX,它容易陷入语法细节。如果让它先写一段 Python 或者 CUDA 伪代码,再翻译成 PTX,生成质量会明显提高。这个技巧在论文里没有提到,但我在实践中发现很有效。