ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

Triton块级编程:从手写CUDA到编译驱动的算子开发新范式

Triton块级编程:从手写CUDA到编译驱动的算子开发新范式 1. 十年坐标从手写CUDA到块级DSL算子开发方式的三次换挡1.1 CUDA时代kernel是一门手艺性能靠经验积累慢慢喂十年前你要让我写一个高性能GPU算子流程基本是这样的打开CUDA先把grid、block、thread想清楚再手动安排每个线程处理哪些数据然后绞尽脑汁处理shared memory、做同步、避免bank conflict、尽量让访存合并。一个稍微复杂的softmax或者矩阵乘从思路成型到性能达标往往要折腾好几天。那个年代算子开发非常依赖个人经验。不是说CUDA写不出来而是写出来和写快之间隔着一条巨大的鸿沟。我记得当年调一个elementwise的kernel为了把global memory的读写换成vectorized load把float4类型安排得明明白白性能提升了将近一倍。这种优化手段说穿了就是你是否熟悉硬件的行为细节。如果把那个时期的算子开发方式做个总结我会说程序员被迫站在硬件的角度看问题。你要想的不是这个计算是什么而是这个计算在硬件上该怎么被切分、怎么被搬运、怎么被隐藏延迟。GPU的能力很强但驾驭它的门槛也很高这也是为什么很长一段时间里能写CUDA kernel的人在团队里都算稀缺资源。1.2 DSL与编译器尝试声明式描述一度很热闹但落地远没想象中顺利大概是五六年前行业里冒出一波用编译器替代手工优化的尝试。Halide提出把计算描述和调度策略分开TVM想把神经网络的算子统一到一个IR里做自动优化XLA则在TensorFlow里直接干掉手动写kernel的诉求。这套思路的逻辑很诱人你只要描述清楚算什么剩下的怎么算交给编译器去搜索和生成。我确实用过TVM写过几个算子也参加过XLA相关的技术分享。坦白说方向上是对的但它有一个绕不开的问题抽象层级选得很尴尬。TVM的调度语法虽然比CUDA友好但你仍然要懂绑定了线程、向量化、流水线这些底层概念而且要花不少时间去学它的调度原语。等于说门槛降低了一点但学习路径换了一套没解决根本问题。更麻烦的是这种完整图级别的编译器一旦塌缩到某几个异构硬件代码生成的效率很难追得上资深工程师手写CUDA。你调优半天可能还是差那么几个百分点。于是业内慢慢形成一个共识大规模、静态shape的模型可以靠编译器但灵活的动态shape、特殊算子还是得有人能手动写点东西。所以那几年的实际状态是——编译器很热闹但真正全自动替代手工的场景少之又少。1.3 Triton作为转折点像写NumPy一样写算子编译器接管硬件细节Triton在这条时间线上的位置其实非常有意思。它既不是传统意义上从头写CUDA也不是TVM那套先描述再调度的复杂语法。OpenAI在做它的时候选了一个更聪明的中间粒度块级编程。你写的还是普通的Python函数只不过在里面用triton.language提供的一系列原语来操作块级别的数据。我第一次打开Triton文档的时候印象很深一个softmax kernel代码量也就三四十行读起来几乎和伪代码一样清晰。你不用去管线程id是多少也不用自己算stride只需要声明我当前这个program要处理X个元素然后像NumPy一样做load、compute、store。剩下的所有东西包括线程怎么组织、访存怎么合并、寄存器怎么分配都是编译器帮你完成的。更关键的是Triton的抽象精度卡在了一个很舒服的位置。它放弃了全图自动优化这种宏大目标专注把单个kernel做到足够好然后和PyTorch生态打成一片。等到PyTorch 2.0把Triton作为torch.compile的默认后端之一大量的研究者忽然发现自己用Python写的高性能算子竟然真的能跑出接近手写CUDA的效果。从这一刻起算子开发的方式才算真正换了一代。2. Triton核心抽象为什么块是人与硬件都能接受的编程粒度2.1 人和硬件对算子的需求完全不同要理解Triton的设计得先搞清楚一个矛盾。从人的角度看我希望写算子的时候只关注数据长什么样、怎么变形、怎么规约最好连指针和线程都不出现从硬件角度看GPU就是一台极度依赖并行和局部性的机器它需要知道每一块数据放在哪、由多少线程一起算、中间变量是不是能放进shared memory。这两个需求天然冲突。你给程序员太多自由就容易回到CUDA的复杂度你给编译器太多自由又容易变成TVM那样需要另一套领域知识才能驾驭。Triton给出的答案是别让程序员管线程但让程序员管块。一个块包含多个元素程序员站在块的视角写计算编译器则在块内部完成最终到线程的映射。举个生活化的例子CUDA像是直接给你一张城市地图让你自己规划从A到B怎么避开拥堵Triton则像是告诉你你去这一片区域把那里的快递都处理了至于你开车走哪条路、中间在哪个路口停一下由导航系统替你决定。对于大部分算子开发者来说后者显然更省心。2.2 一个简单的softmax kernel看Triton如何表达块计算我直接贴一个最经典的softmax实现你看一眼就明白我在说什么import torch import triton import triton.language as tl triton.jit def softmax_kernel( output_ptr, input_ptr, input_row_stride, output_row_stride, n_rows, n_cols, BLOCK_SIZE: tl.constexpr, ): row_start tl.program_id(0) col_offsets tl.arange(0, BLOCK_SIZE) input_ptrs input_ptr row_start * input_row_stride col_offsets mask col_offsets n_cols row tl.load(input_ptrs, maskmask, other-float(inf)) row_minus_max row - tl.max(row, axis0) numerator tl.exp(row_minus_max) denominator tl.sum(numerator, axis0) softmax_output numerator / denominator output_ptrs output_ptr row_start * output_row_stride col_offsets tl.store(output_ptrs, softmax_output, maskmask)这段代码里没有threadIdx.x没有blockIdx.x没有shared memory也没有同步屏障。你只需要知道三件事每个program处理一行tl.arange生成列的偏移量tl.max和tl.sum沿着指定的轴做规约。编译器会自动把这一行数据切到足够多的线程上必要时把中间结果塞进shared memory做规约然后再写回去。我第一次跑这个kernel的时候最惊讶的是性能居然完全能打。对比同样配置下手写CUDATriton版本通常只慢个百分之几有些情况下甚至因为自动向量化做得好而反超。这就是块级抽象的价值你保留了这一个块里面的数据我要怎么变换的控制权却把块内怎么组织执行交给了编译器。2.3 编译器在块之上做的事访存合并、shared memory与调度有人会问编译器凭什么能做得比人更好答案是它在一个非常狭窄但关键的范围内做deep optimization。Triton的编译器拿到块级操作后会分析每一段访存的shape、步长、对齐方式然后自动生成合并访问模式尽量让相邻线程访问相邻地址提高global memory带宽利用率。这一点和手写CUDA里最重要的优化思路完全对齐只不过由编译器推理完成。同理当你做softmax这种需要跨元素规约的操作时编译器会识别出块内归约模式自动插入shared memory分配和同步逻辑。你在Python代码里一行同步都没写但生成的PTX里该有的bar.sync一条不少。这种能力背后就是LLVM/MLIR那一整套现代编译器基础设施它会在Triton IR层自我发现哪些块操作可以映射成硬件原语哪些必须拆成更小的步骤。这其实就是很多人提到的算子自发现不是人去找算子的可优化机会而是编译器在IR层面系统性地扫描访存模式、计算模式、复用距离然后自动套用优化规则。Triton的聪明之处在于它把搜索空间限制在块内避免了TVM那样在全图上暴力搜索调度方案所以编译速度快、生成代码质量也稳定。3. 一个Triton算子从Python到GPU执行的全流程3.1 jit装饰与AST捕获Python函数如何变成编译器IR很多人第一次看Triton会以为triton.jit装饰过的函数还是普通Python函数。实际上当你用指定的指针参数去调用它时Triton并不会真的逐行执行Python而是捕获这个函数的AST做类型推导然后翻译到自己的中间表示Triton IR。这个过程有点像把Python语法糖拆掉还原成一个数据流图。循环、if、算术操作都会被转成静态的、带形状信息的中间表示。副作用就是Python里很多动态特性比如任意对象方法、闭包捕获复杂对象在kernel里是不能用的。能用的是tl.arange、tl.load、tl.store这些建立在形状推断基础上的核心原语以及普通整数、浮点运算。当初写的时候我还不习惯老想在kernel里调用Python库函数结果编译直接报错。后来才明白Triton要求你以块为单位思考你写的是一段会被编译器整体分析和变换的代码而不是一行行解释执行的脚本。这个心智模型转过来之后很多东西就顺了。3.2 Triton IR到LLVM/PTX自动向量化、寄存器分配与代码生成AST捕获完接下来的链路是Triton IR - MLIR/LLVM方言 - LLVM IR - PTX - SASS最终由NVIDIA驱动做JIT编译。很多人把Triton简单理解成Python到PTX的翻译器其实中间的优化多得很。在Triton IR层编译器会根据块的形状和访存对齐信息做向量化。如果你处理的是连续的float数据它会优先生成128-bit的向量load如果对齐条件不满足则回退到更窄的访问。这一层的决策直接决定你最终能不能吃满显存带宽。接着是shared memory的liveness分析编译器会安排中间缓冲区减少global memory的反复读写。到了LLVM阶段常规的循环不变量外提、指令合并、寄存器分配全部生效。这也是为什么Triton生成的kernel性能往往不弱于手写CUDA的核心原因之一它同时吃到了高层块级语义带来的精准优化机会和底层LLVM十几年打磨的优化能力。我做过一次实验同一个GELU算子Triton编译出来的PTX里面居然有自动生成的__nv_bfloat16类型转换指令比我手写CUDA时更懂得利用硬件特性。3.3 launch与硬件调度从grid循环到SM上的执行器代码生成完毕剩下的就是launch了。Triton在Python端拿到你传入的grid大小为每个program instance生成对应的起始地址和参数组合。GPU的调度器会把这些block按顺序分发到流多处理器上每个SM内部再按照warp调度器把线程束放行。这里有个容易忽略的点当你的grid特别大时一个program instance并不是独享一整块连续的执行时间。SM会在空闲时插入别的block形成一种硬件级别的overlap。Triton的编译器其实无法直接控制这个层面但它会尽量让每个block内部有足够的独立工作量用更多并行度来掩盖访存延迟。这跟手写CUDA的道理一样只不过你在Triton侧通过调整num_warps间接影响block的线程规模而不是手写一切。配合torch.cuda上的launch机制整个过程在Python侧看起来就是一个普通的函数调用。你要是只关心功能完全不用管中间发生了什么但如果你后面要排查性能问题就必须要能顺着Python wrapper - Triton kernel - grid设置 - 实际SM执行这条链路去定位瓶颈。3.4 验证与性能分析别让能跑骗了你跑通一个Triton算子不算难事难的是确认它真的对。我见过很多新手在对比kernel输出和PyTorch参考实现时直接assert结果遇到float精度问题一比对就崩然后怀疑Triton代码写错了。正确做法是先用torch.allclose设定合适的atol和rtol特别是涉及exp、reduce这类容易累积误差的操作。性能分析也不要只盯time.time()。Triton自带的triton.testing.do_bench比手动计时靠谱得多它会做warmup自动算平均和标准差。更进一步你可以在Nsight Compute里看到Triton生成的kernel名字然后逐个看访存吞吐、计算吞吐和warp occupancy。我第一次用Nsight看Triton的kernel时还特意找了一下名字确认它确实是被正常编译和执行的而不是走了什么magic路径。事实证明它就是普通kernel只是生成过程自动化了而已。4. 大量算子叠加真正挑战从来不是单个kernel4.1 单算子优化和整图性能之间的鸿沟很多做算子开发的人一开始都会陷入一个误区把单个kernel优化到极致模型整体性能就该很好了。但当你真正去分析一个PyTorch模型时会发现里面可能有几百个算子。其中很多算子本身很简单比如加一个bias、过一个激活、做一次layout转换单独看每个也就几十微秒但架不住数量多。每个kernel的launch overhead、参数传递、两次kernel之间的全局内存写回都会成为性能的黑洞。更麻烦的是kernel和kernel之间的数据来回倒腾往往会把本来能驻留在寄存器或shared memory里的数据逼到global memory里走一遭白花带宽。这其实是大量使用算子对硬件性能的挑战最本质的部分你优化的不是数学运算而是数据搬运的路径。我看过一个很典型的模型trace里面光elementwise类型的算子就有三十多个。如果每个都单独跑GPU的算力利用率可能连10%都不到。这就是为什么无论编译器怎么演进算子融合都是绕不开的话题。不做融合堆再多的优化技巧都是在给搬运工加工资。4.2 访存瓶颈与算子融合为什么matmulprelu这类融合成为标配当模型里出现大算子和小算子交替时最划算的做法就是让小算子贴在大算子身上让中间结果留在片上。一个非常典型的例子就是matmul prelu矩阵乘法的输出如果先写回global memory再由下一个kernel读出来做prelu等于多了一整轮DDR流量如果直接在矩阵乘的epilogue阶段把prelu应用在寄存器里的结果上访存开销几乎可以忽略。这类融合算子之所以越来越流行是因为它不仅出现在GPU上在NPU上也一样重要。我查过一些昇腾上的算子优化案例AscendC里做matmulprelu融合的思路和CUDA上做FMA融合本质上没有区别把矩阵乘的累加结果留在AI Core内部紧接着做激活最后一次性写回。换的是硬件名字不变的是减少片外搬运这条原则。在Triton里做这种融合也非常自然。你可以在load完一个tile、做完tile级别的矩阵累加后直接对寄存器中的结果做tl.where之类的激活操作然后再store。编译器会帮你把中间状态锁定在片上不会额外分配大块global memory。我写过一个带LeakyReLU的matmul融合后整体算子比分开跑快了将近30%而且代码只多了三行。4.3 Triton实战调优num_warps、num_stages和autotune的取舍到了调优环节Triton提供了几个核心旋钮num_warps、num_stages、以及块大小。很多人一上来就照着默认值跑性能不太行就放弃了其实这几个参数对性能影响非常大。num_warps控制一个program内部的线程束数量直接影响并行度。块太小、元素太少时开太多warp反而会因为同步开销增大而变慢块很大、计算很重时warp太少又不足以掩盖访存延迟。num_stages用于软件流水线在load下一块数据的同时计算当前块用shared memory做缓冲隐藏global memory访问延迟。这个参数对访存密集算子尤其关键。我通常的习惯是把关键算子交给triton.autotune让它在少数几组配置里自动搜索。别搞几百组冷启动会怀疑人生4到8组配置足够了。我实测下来num_warps4配合num_stages3对很多elementwise和softmax算子都是不错的起点而matmul类更偏向更大的块和更深的流水线。4.4 编译时间、缓存与动态shape工程化中最容易翻车的地方Triton虽然运行快但编译一个kernel不是完全零成本的。每次你改变launch配置、块大小、数据类型、甚至某些tl.constexpr参数的值编译器都可能要重新走一遍优化链路。在训练循环里频繁换shape最常见的结果就是卡顿一下然后速度正常那一卡就是编译的开销。Triton本身有kernel缓存默认会按源码、参数、后端等信息做哈希落到磁盘上所以同样配置重复执行不会重复编译。但动态shape会让缓存命中率直线下降。如果某个维度在训练中频繁变化我建议干脆把它当成BLOCK_SIZE的constexpr处理时设置一个覆盖所有shape的公共块大小避免为每一个shape都生成一份kernel或者用padding把形状对齐到同一档位收益远大于浪费的那点计算。还有个小坑是环境里的TRITON_CACHE_DIR默认目录空间如果被占满编译可能会失败。离线部署场景我一般会显式指定一个可写的缓存路径并在镜像里预编译一遍关键算子把缓存一起打包进去这样推理服务器起来后就不用现场编译了。5. 从CUDA到NPUTriton如何应对多硬件时代的算子开发5.1 NPU带来的新变量AI Core、带宽层级与调度方式过去做高性能算子开发大家默认目标就是NVIDIA GPU。但近两年NPU的声音越来越大尤其是昇腾这类AI芯片核心计算单元叫AI Core很多设计逻辑和GPU的SM类似但又不完全一样。你在CUDA上习惯的那些假设——比如统一的shared memory大小、固定的warp宽度、L2缓存的分配方式——在NPU上可能全部要重设。NPU上算子开发的难点不完全在算力不够而在搬运和调度的自由度不同。有些NPU对局部内存的容量和Bank组织有自己的讲究对多核任务切分也有一套固定的同步原语。直接在AscendC里写算子有点像当年手写CUDA性能天花板高但开发链路长、调试工具相对少。这时候Triton这类高级抽象就有机会了它把块级计算变成一种硬件无关的描述再由不同的后端映射到具体架构。5.2 MLIR/LLVM作为中间桥梁Triton IR映射到AscendC等后端Triton能被多个硬件后端盯上核心原因在于它的IR设计比较干净没有过早绑定PTX。社区里其实已经有很多人在做Triton到非NVIDIA后端的适配基本路径是先复用Python前端和Triton IR然后把块级操作lower到MLIR方言再由各硬件厂商自己的编译器后端继续做调度和代码生成。对应到昇腾生态你会看到两条路线一条是直接用AscendC手写融合算子面向量产和极致性能另一条是把Triton kernel翻译成类似的IR再映射到AscendC的C代码结构上。对我来说Triton这条路线更大的价值在于快速原型验证我先用Triton把算法逻辑写清楚跑通shape和数值再决定要不要花力气用AscendC手工精调。当然能编译不等于高效。不同NPU的向量单元宽度、矩阵单元形态、内存层级差异很大Triton的自动优化规则在GPU上学到的那套经验直接搬过去不一定合适。所以多后端适配的现实状态是前端通用后端逐硬件精调。这也是为什么各家在做落地时都不会只依赖Triton一条链而是把Triton当成易于生成、易于验证的入口。5.3 多后端适配的现状与坑能编译不等于高效真实踩过坑的人会告诉你跑通NPU后端和跑出高性能之间隔着一个巨大的调优深渊。跨后端之后原本在GPU上不需要关心的细节会突然冒出来NPU的core之间通信方式可能不是全局同步block的切分方式会影响片上数据复用甚至不同shape的padding策略都会导致计算效率翻倍变化。另外Triton社区的PR和更新基本还是以NVIDIA backend为主非NVIDIA后端的维护活跃度、bug修复速度都没法和CUDA路径比。如果你所在团队要长期做某款NPU的算子库我建议不要把整个技术栈押在社区维护的第三方后端上而是把它作为辅助验证工具核心算子还是要有自己的性能回退方案。从更长远的视角看算子开发正在走向一套块级语义多个硬件后端的格局。将来一个模型要在GPU、NPU、甚至CPU上跑开发者需要的不是为每个硬件单独写一遍算子而是把算子的逻辑描述好然后让编译器针对目标硬件生成最合适的执行方案。Triton在这方面至少开了一个很实际的头。6. 实验记录搭建Triton环境、跑通第一个算子并避坑6.1 安装与版本匹配pip一条龙没那么简单很多人上来就pip install triton装完跑demo大概率会撞上版本和PyTorch对不上的问题。Triton和PyTorch的CUDA runtime走得比较近版本错位时经常出现Illegal instruction或者奇怪的CUDA error。我现在的做法是优先使用PyTorch自带匹配的Triton而不是自己额外装一份最新的。如果你用pip安装PyTorch官方wheel它通常已经捆绑了一个兼容版本的Triton需要单独升级Triton时再去triton的release页面找对应的wheel包。源码编译也是一种方式但耗时较长还要装好LLVM依赖除非你想改编译器的源码做二次开发否则不建议一上来就跳进源码编译的坑。装完可以跑一个最简单的加法kernel验证环境。重点看它能不能正常生成PTX以及triton.testing.do_bench能不能正常计时。如果能跑通说明CUDA环境、Python版本、编译链都OK了。6.2 从vector add到softmax上手只需要两个晚上学习Triton最快路径我建议直接从vector add开始别看太多入门教程。第一晚写一个vector_add_kernel传入两个指针和大小用tl.arange和mask处理边界跑通和PyTorch结果一致。第二晚改写成softmax加上tl.max和tl.sum做行规约顺便试一下num_warps4和num_warps8的性能差异。这两个例子足够覆盖80%的常用语法load/store、mask、arange、reduce、constexpr。之后再接触矩阵乘法的tl.dot、原子操作的tl.atomic_add、多级流水的num_stages都会非常快。我甚至见过一个完全没碰过CUDA的研究者用两周时间把一套Transformer里的自定义算子全部换成Triton实现跑出来的性能和原来手写CUDA版本基本持平。这放在十年前是不可想象的。6.3 我踩过的三个坑版本兼容、边界shape、自动调优冷启动坑一版本兼容。有一回我把Triton升到大版本结果旧代码里tl.load的eviction_policy参数行为变了编译不报错但性能掉了20%。排查了半天才定位到是这个语义变化。所以升级Triton之后一定要对着release note核对一遍你用了哪些高级参数别默认它们是稳定的。坑二边界shape。block size和实际shape不一致时mask必须有。但mask太多会影响访存效率因为编译器没办法确定哪些地址一定能合并访问。我的经验是尽量让shape对齐到block size的整数倍实在不行再开mask。比如做softmax时把一行元素padding到2的幂经常比用精确mask跑得更快。坑三自动调优冷启动。autotune配置太多组程序第一次启动会逐个编译候选kernel几十个组合下来可能等了好几分钟才正式开始跑。这在本地开发还能忍在线上服务或大规模评测时会非常致命。我现在都严格控制候选集数量并把缓存目录稳定下来如果benchmark的场景很固定还不如直接手动选一组最合适的配置把编译时间压到一次。最后再分享一点我的个人感受做算子开发这十年最明显的变化是写算子的门槛在快速下降但判断一个好算子需要什么的能力反而更值钱了。Triton帮你省掉了手工管理线程的体力活可要写出能在不同硬件上都稳定高效的算子你还是得理解访存、融合、调度这些底层逻辑。工具在迭代基本功不会过时。如果你正准备入坑算子开发我的建议就是别怕底层的那些概念先上手写几个Triton算子再回过头去看CUDA你会发现很多曾经需要死记硬背的优化细节现在都变得顺理成章了。
返回列表