
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), dtypefloat16), B: T.Tensor((K, N), dtypefloat16), C: T.Tensor((M, N), dtypefloat16), ): for m, n in T.Parallel(M // BM, N // BN): acc T.alloc_fragment((BM, BN), dtypefloat32) T.clear(acc) for k in T.serial(K // BK): with T.magma_tile(m, n, k): A_tile T.alloc_shared((BM, BK), dtypefloat16) B_tile T.alloc_shared((BN, BK), dtypefloat16) 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), dtypefloat32)分配的是一个寄存器片段。注意这里用的是 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), dtypetorch.float16, devicecuda) A torch.randn((M, K), dtypetorch.float16, devicecuda) B torch.randn((K, N), dtypetorch.float16, devicecuda) 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 启动参数非法。经验性的起点是BM128, BN128, BK32这是很多公开例子的默认值在 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”的判断力。