☰
tilelang:GPU算子的全局融合之道,从原理到调优实践
2026/10/8 23:31:07 网站建设 项目流程

tilelang 这个名字,最近在 GPU 优化圈子里出现频率相当高。很多人拿它和 Triton 放在一起聊,说它擅长“全局融合”,能在 GEMM、FlashAttention 这类算子上超过手写 CUDA 的性能。作为一个在 Triton、CUTLASS 和手写 CUDA 之间来回折腾过的人,我花了一段时间把 tilelang 用进实际项目,也踩了些文档里没写清楚的坑。这篇文章不打算复述官方文档,而是想聊清楚:tilelang 到底在解决什么问题,它的核心设计思路是什么,我在集成和调参过程中吃了哪些亏,以及什么样的团队可以认真考虑把它接进生产链路。

1. 先搞清楚:tilelang 要替代的是什么

1.1 过去写高性能算子的三种套路,各自卡在哪里

做 GPU 算子开发的人,手上基本逃不开三样东西:手写 CUDA、依赖厂商库、或者用 DSL 编译器。

手写 CUDA 是性能上限最高的路线,你对线程块、共享内存、寄存器、访存模式都有完全控制,CUDA 编程模型能表达的东西非常底层。代价也很明显:一个稍微复杂一点的算子,从写核函数到调优,少则几天,多则几周。特别是当你处理的是“融合算子”——比如把 GEMM、激活、LayerNorm、残差连接都串在一个 kernel 里——涉及的循环合并、共享内存分配、同步策略、寄存器压力管理,每一项都是高手才能玩得转的东西。团队里能稳定产出这种代码的人,一般都比 GPU 还难找。

调用 cuBLAS、cuDNN 这类厂商库是最省事的路线。性能通常不差,因为厂商针对目标硬件做了大量手工优化。但问题在于它是黑盒。你没法把它的底层算子跟自己的业务逻辑做深层次融合,也没法针对某个特殊 shape 或特殊数据布局做裁剪。遇到 cuBLAS 表现一般的长条形 GEMM、超大 batch 的小矩阵、或者 attention 这类厂商库还没覆盖成熟的新算子,就只能干着急。

于是很多人转向 Triton、TVM 这类 DSL。Triton 的抽象很好,把程序员从线程层面解放出来,让你以 tile 为最小单位思考,编译器负责把 tile 映射到线程和寄存器上。但 Triton 有一个绕不过去的短板:它擅长生成单个 kernel,却不太擅长跨 kernel 做全局优化。你在 PyTorch 里写的算子序列,通常还是要拆成多个 Triton kernel 或者混合 PyTorch 原生算子,中间结果反复写回显存,带宽和时间都在这些来回倒腾里浪费掉了。

1.2 tilelang 的出现,正好堵在这个空白上

tilelang 的核心目标,用一个词概括就是“整图融合”。它想让用户先写清楚完整的计算逻辑,然后编译器基于某种 tile 抽象,把整个计算过程编译成一个整体的大 kernel,而不是由用户手动拆成一段一段去优化。

我最初看到这个思路时,第一反应是:这不就是编译器领域的老话题吗,析构、融合、流水线,TVM 不是早就做过?后来实际用下来才理解,tilelang 和 TVM 有一个关键区别——它把 tile 这个计算块提到了绝对核心的位置。它不是在一个通用 IR 上面东补一块西补一块,而是从描述语言的第一天起,就要求你以 tile 为单元表达计算,编译器再从 tile 之间的依赖、访存、调度关系里挖掘性能空间。这有点像写文章和做拼贴画的关系:一个是先写好句子再调结构,一个是先从整体版式出发让素材自动归位。

这样一来,编译器能做的事情就比“把一个循环展开”深刻得多。它可以跨 tile 合并循环,可以自动决定哪些数据放共享内存、哪些放寄存器,可以自动编排流水线和 double buffering,还可以把多级 memory hierarchy 的搬运路径在整个 kernel 范围内统筹规划。最后你用掉的开发时间接近写 Triton,但生成的代码在访存局部性、指令级并行度上,又向手写 CUDA 靠拢了一大步。

2. 核心设计拆解:为什么 tile 能成为编译器的“主语”

2.1 MagmaTile 到底是个什么东西

第一次打开 tilelang 文档的人,十有八九会被magma_tile这个词唬住。我刚开始也以为这是什么玄学概念,后来才有点感觉:它想表达的可能是“一块凝固的计算流”。

在 tilelang 里,tile 不是一个简单的二维分块,它携带的信息比“矩阵切一小块”多得多。一个 tile 本身带有循环结构、访存模式、数据类型和数据依赖关系。你可以理解为它是一段“带形状的程序”,编译器拿到这段程序后,可以做跨循环合并,可以决定数据放在哪一级存储,可以自动把 tile 之间的依赖关系变成流水线。最终生成的 SASS 或者 PTX 指令里,tile 通过编译器的手被拆成线程束级别的操作,但你在源代码层根本不需要碰这些细节。

这种抽象方式最大的好处是,它给了编译器一个“稍大但足够规整”的分析单元。如果编译器直接面对的是任意嵌套循环和到处乱飞的 pointer,它做融合和资源分配的难度很大;如果只面对向量加法这种单元素操作,它能发挥的优化空间又太小。tile 恰好卡在中间:够大,大到能看出全局的访存和计算模式;够规整,规整到编译器可以用确定的规则去调度。

2.2 编译时的三个关键动作:融合、搬运、调度

把整段计算编译成一个 kernel,涉及三件核心的事。

第一件事是融合。tilelang 会把计算图里的相邻算子合并到一起,消除中间张量在显存里的写回和读取。这个优化在 PyTorch 里做一次torch.jit.script或torch.compile也能得到一部分,但 tilelang 做得更彻底。它不只是在图级别把相邻 Op 合并,而是深入到 tile 内部,把 GEMM 的累加、偏置的加法、激活函数的计算、甚至后续的归一化,全部糅合到寄存器和共享内存的片段里,中间结果根本不会离开芯片。

第二件事是搬运。现代 GPU 的存储层级相当复杂:显存、L2、共享内存、寄存器,每一级的带宽和延迟差一个数量级。如果搬运路径安排得不好,再快的计算指令也会被访存拖死。tilelang 的编译器会分析每个数据片段被使用的时间窗口,决定它应该在哪一层存储里待着,什么时候从显存搬到共享内存,什么时候从共享内存展开到寄存器,什么时候可以丢掉。这套决策逻辑本质上是在跟延迟做博弈——你把数据搬早了,占着宝贵的共享内存不能用;搬晚了,计算单元饥荒,流水线空转。编译器需要找到那个利益最大化的时机。

第三件事是调度。同一个 tile 计算,可以映射成不同的线程束数量、不同的数据切分方式、不同的循环顺序,性能差个两三倍是非常正常的事情。tilelang 内置了自动调度机制,会在这套空间里搜索一个合理配置。它不需要用户去算“这个矩阵块应该分给多少个线程”,只要用户指定好 tile 的形状,剩下的编排放置由编译器接手。

2.3 这种设计到底带来了什么实打实的收益

我跟进过几个用 tilelang 重构的算子,最直观的感受是 kernel 数量显著变少。以前一个复杂的 attention 模块,可能要拆成 5 个甚至 8 个 kernel,分别负责 QK 乘法、softmax、PV 乘法、输出投影、残差相加。用 tilelang 重写之后,一个融合 kernel 就把事情全干了。kernel 少了,带来两个连锁好处:一是 GPU 不需要频繁启停,kernel launch 的开销被抹掉;二是数据不用反复在显存和计算单元之间兜圈子,整体的 memory-bound 特征被大幅改善。

更重要的一点是,tilelang 的编译器让“从研究到落地”的路程变短了。以前想尝试一个新的融合策略,要先画图分析数据流,再盯寄存器分配,再上 Nsight 看 stall 原因,三四个来回下来热情已经被消耗完。现在改的是描述层的逻辑,编译器替你处理底层的脏活。当然,这不是说完全不需要底层知识——如果你连共享内存和 bank 冲突都不了解,那调参的时候依然会一头雾水。

3. 实操手记:把第一个 tilelang 内核跑起来

3.1 安装与环境对齐,这一步比想象中重要

tilelang 支持 pip 安装,命令很简单:

pip install tilelang

但安装本身不是难点,环境对齐才是。我一开始在 conda 环境里直接装,结果发现跟 torch 的 CUDA 版本对不上,编译出的 kernel 要么跑不了,要么报错提示得很含糊。后来总结出一个稳妥的做法:先建一个干净的 conda 环境,从官方源装好匹配的 PyTorch,再装 tilelang。版本上我实测比较稳的是 CUDA 11.8 或 12.1 配 PyTorch 2.1 以上。别在生产环境里直接升级 PyTorch 来将就 tilelang,否则那些底层二进制可能全军覆没。

注意:tilelang 的 API 迭代速度很快,我下面写的代码是“理念级”示例,不保证和你安装的版本逐字对应。上手前先把本地的tilelang.__version__打印出来,再看对应的文档或示例代码,这样能少走很多弯路。

3.2 一个最小可用的 GEMM 内核长什么样

GEMM 是 GPU 生态里的“hello world”,用来建立对 tilelang 的心智模型最合适。核心逻辑是:声明输入和输出张量的形状,用 tile 描述循环和数据结构,然后让编译器自行处理底层实现。

import tilelang import tilelang.language as T M, N, K = 2048, 2048, 2048 BM, BN, BK = 128, 128, 32 @T.prim_func def gemm( A: T.Tensor((M, K), dtype="float16"), B: T.Tensor((K, N), dtype="float16"), C: T.Tensor((M, N), dtype="float16"), ): for m, n in T.Parallel(M // BM, N // BN): acc = T.alloc_fragment((BM, BN), dtype="float32") T.clear(acc) for k in T.serial(K // BK): with T.magma_tile(m, n, k): A_tile = T.alloc_shared((BM, BK), dtype="float16") B_tile = T.alloc_shared((BN, BK), dtype="float16") T.copy(A[m * BM, k * BK], A_tile) T.copy(B[k * BK, n * BN], B_tile) T.gemm(A_tile, B_tile, acc) T.copy(acc, C[m * BM, n * BN])

这份代码里有几个值得琢磨的点。

T.Parallel和T.serial的区别是学习 tilelang 的第一道坎。T.Parallel表示循环的迭代之间是并行关系,编译器会把每一次迭代分配到一个独立的计算块,它们天然适合并行执行。T.serial则表示依赖顺序很重要,必须按照 k 的次序依次执行——在 GEMM 里这是累积的过程,每个 k 步的乘法结果都要加到同一个累加器上。如果你把T.serial错写成更激进的并行循环,结果很可能是错的。

T.alloc_fragment((BM, BN), dtype="float32")分配的是一个寄存器片段。注意这里用的是 float32,不是 float16。原因是 GEMM 累加过程中会产生大量浮点误差,而累加器用更高精度保存是通用工程实践。我曾经在图省事的时候把累加器直接声明成 float16,结果算出来的结果误差大到不可接受。这一条原则不只在 tilelang 里成立,在任何 GEMM 实现里都应该记住。

T.copy的行为也有讲究。这个 copy 不是简单的赋值,它会编译成经过深思熟虑的访存指令,涉及合并访问、对齐、甚至双缓冲的准备。从显存到共享内存的拷贝、从共享内存到寄存器片段的展开,都是通过T.copy表达的。你不需要显式写同步语句,编译器会自动在合适的位置插入屏障或者做软流水。

最后,T.gemm是一个语义化调用。它告诉编译器“这里有一块矩阵乘需要计算”,具体怎么切分线程、怎么利用张量核心、怎么处理尾数,全部交给编译器。这也是 tilelang 和手写 CUDA 最大的体验差异——你把意图说清楚,把资源约束摆明白,实现细节由编译器给出。

3.3 把编译好的内核接进 PyTorch

编译和调用也很直白,大体是这样的模式:

kernel = tilelang.compile(gemm) C = torch.empty((M, N), dtype=torch.float16, device="cuda") A = torch.randn((M, K), dtype=torch.float16, device="cuda") B = torch.randn((K, N), dtype=torch.float16, device="cuda") kernel(A, B, C) torch.cuda.synchronize()

这里有一个和常规 PyTorch 不同的地方:tilelang 编译出来的 kernel 是同步阻塞式的还是异步的,取决于底层实现和当前设置。我建议你在正式做 benchmark 之前,统一加一次torch.cuda.synchronize(),避免被 kernel 执行队列的异步性欺骗,测出虚高的时间。

如果要把 kernel 包装成一个可导的自定义算子,思路也很清晰:写一个torch.autograd.Function,前向里调用 kernel,反向里再调用你自己实现的配套反向 kernel。反向的 kernel 不一定非要用 tilelang 写,但如果你追求端到端性能,最好也顺手用 tilelang 把反向融合算子写了,这个项目在这个方向上的便利性是很明显的。

3.4 第一次调参时该往哪个方向使劲

跑通之后,很快会进入调参环节。我最先做的尝试是调整BM/BN/BK这三个分块大小。这个选择会直接影响共享内存占用、寄存器占用和访存模式。太小,会让数据复用性不足,访存开销占比上升;太大,会让单块工作量过高,或者直接溢出共享内存,导致编译失败或者 kernel 启动参数非法。经验性的起点是BM=128, BN=128, BK=32,这是很多公开例子的默认值,在 A100 和 H100 上都比较稳。

第二件要关注的是BK的选择。BK决定每个 k 步搬运多大数据进共享内存,它的背后是对访存带宽和计算密度的权衡。如果BK偏小,双缓冲的优势还没发挥出来就切换到下一步;如果偏大,共享内存压力上升,能驻留的并发块数量下降。实际测试时,可以以 2 的幂次从 16 试到 64,观察吞吐和显存占用变化,通常能找到明确拐点。

第三件值得尝试的是把循环次序从串行改成带流水线。T.serial是稳妥的串行循环,但现代 GPU 靠并行掩盖延迟,单纯的串行循环会浪费很多算力。tilelang 的文档里经常出现流水线相关的配置项,开启后编译器会在循环迭代之间做双缓冲和预取,把访存延迟藏进计算时间。这一步经常是性能从“还不错”提升到“接近手册上限”的关键。

4. 选型视角:tilelang、Triton、手写 CUDA 各有各的主场

4.1 三种方案的对照表

维度tilelangTriton手写 CUDA
抽象层级tile/程序块tile/程序块线程/线程束/程序块
全局融合能力强,设计目标就是整图融合一般,偏向单内核优化完全可控,但成本极高
开发效率高,接近写 Python 数学公式高低,需要大量底层排错
性能代表性高,接近手工实现中等偏上高,取决于作者水平
学习曲线中等,需要理解 tile 抽象较低,上手快陡峭,掌握硬件细节
生态成熟度还在快速迭代,文档变化快相对成熟,社区大稳定
适用场景需要融合算子的推理/训练优化快速原型和通用算子开发旗舰算子和极致性能调优

这个表里最关键的一行是“全局融合能力”。Triton 虽然也支持在单个 kernel 里做不少事,但当你把注意力、归一化、残差连成一个复杂计算流时,Triton 通常需要借助 PyTorch 的图编译器在多个 kernel 之间协调,或者由用户手动把它们掰成一块。tilelang 从描述层就鼓励你把整个计算流写在一个 block 里,融合是默认动作,不是可选项。

4.2 手写 CUDA 还有没有存在的必要

有,而且长期会有。tilelang 再怎么自动调度,也不可能覆盖所有硬件架构的特性和所有算子的特殊形态。当你面对一个极不规则的访存模式,或者需要逐指令抠 SASS 级别的细节时,手写 CUDA 依然是终极兜底方案。

我用 tilelang 过程中很明显的感受是:它把“通用优化”这块做得出色,但“特殊优化”还是要靠人。比如某个算子的输入有非常特殊的矩阵结构(带状矩阵、三角矩阵),或者你需要做多核之间的精细协同,这类知识没法完全灌输给编译器。所以我的建议是,团队里最好还是有人能读懂 PTX/SASS,能在 tilelang 生成的代码不理想时清楚地指出瓶颈位置,而不是盲目调参。

4.3 什么时候我建议直接用 tilelang

判断标准可以很简单:你的性能瓶颈是不是来自“算子之间的来回搬运”和“多 kernel 启动开销”。如果是,这就是 tilelang 的主场。以 FlashAttention 为代表的融合算子是最典型的一类——它天然要求把多个步骤揉进一个 kernel,传统的拆开实现性能损失非常明显。tilelang 的公开示例里就包含大量这类算子,很多实测结果比 manually optimized 的版本还要快。

如果你的项目里布满了各种自定义的、组合式的推理算子,并且团队里有懂 GPU 底层的人,完全可以考虑把 tilelang 加进技术栈。它不会取代你手里的所有工具,但会在“既要快速交付又要高性能”的矛盾点上,给你提供一个很舒服的中间选项。

5. 常见问题与排查技巧实录

5.1 编译时间很长,甚至感觉卡死了

我第一次编译稍微复杂的 attention 融合算子时,等了将近一分钟,还以为进程挂了。后来才知道,编译器在自动调度阶段要进行大量搜索,搜索空间跟 tile 的形状、循环嵌套深度直接相关。这不是死机,是它在试着找到最优配置。

如果你的编译时间实在不可接受,可以先把调度搜索的范围调小,或者关闭一些高级优化 pass,先把正确性验证通过,再逐步打开高级优化。另外一个实用技巧是,把编译结果缓存起来,避免每次跑代码都重新编译。对于一天内反复迭代的实验场景,这个缓存的收益极大。

5.2 shared memory 爆了,或者 bank conflict 莫名其妙

当你把分块调大时,最先遇到的就是共享内存溢出。工具会报出具体数字,你一看就知道是超出了硬件额度。解决方向是缩小分块,或者改变数据布局,减少冗余。我们曾经因为把整个矩阵块原封不动搬进共享内存,而没做 swizzle 处理,导致 bank conflict 严重,吞吐直接掉了一半。如果你在分析工具里看到很高的 shared memory bank conflict 计数,第一反应应该是检查数据的存放顺序,而不是急着调 block 大小。

5.3 接入 PyTorch 自动求导时的反向缺失

tilelang 给的只是 kernel 调用方式,不会替你写反向传播。你需要自己实现反向的融合 kernel,然后在torch.autograd.Function里把它们串起来。这个坑我踩过两次——第一次以为 forward 能跑通就万事大吉,结果一训练就报grad相关错误;第二次则是反向 kernel 没考虑清楚,数值对不上。建议是先用 PyTorch 的自动微分结果做一遍数值对照测试,确认相对误差在 1e-3 量级以内,再放进训练流程。

5.4 如何确认编译器真的生成了好代码

不要只看time命令的报告,要亲眼看到生成的代码长什么样。tilelang 提供了生成 kernel 源码的方法——你可以把生成的 CUDA 源码或 PTX 导出来看,也可以配合 Nsight Compute 跑 profiling。我习惯先用ncu看几个关键指标:计算吞吐、访存吞吐、stall 原因分布。如果发现long scoreboard占比很高,说明访存延迟没被隐藏,优先考虑双缓冲和预取;如果发现barrier重,说明同步开销是瓶颈,看看能不能减少显式同步或者调整 tile 划分。

一个我反复使用的检查清单很简单:先跑一次小规模 shape 验证正确性,再跑一次中大规模看显存占用趋势,最后用 profiling 工具定位瓶颈。任何一步不对劲,都别急着上生产环境。

6. 最后分享几个我从实践中得到的体会

如果让我只保留一条建议,那就是:先把 tilelang 当成一个“研究型工具”来用,而不是第二天就要上生产线的万能钥匙。它的确很强,但还处在快速迭代期,API 会变,编译策略会变,周边生态也会变。你在这个阶段投入学习,收获的是对 GPU 编译优化更深的理解,而不只是会调一个框架的接口。

我在实际项目里养成的一个习惯是,每个用 tilelang 写的 kernel 都配一个 Triton 或 PyTorch 原生实现的 baseline。优点是双重保障:一是每次改动都有明确的对比基准,不会自我感觉良好;二是如果 tilelang 后续版本某个行为发生变化,我能立刻察觉。性能调优这件事,最怕的就是凭感觉,有一个可复现的对照实验,比任何玄学的“我觉得它会快”都靠谱。

还有一个小技巧是,多去看官方示例里那些 kernel 的写法,尤其是融合 attention 类的例子。这类例子几乎把 tilelang 的精华全展示出来了——如何组织 tile、如何管理共享内存、如何切分循环。把它们读懂、自己改着跑一遍,你对这个项目的能力边界会有一个比文档描述清晰得多的认识。我自己就是从这个过程中,逐步建立起了“哪些算子适合交给 tilelang,哪些还得自己写 CUDA”的判断力。

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

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

立即咨询