ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

手写Triton Kernel跑通Qwen3.5-0.8B:21算子拆解与CUDA Graph优化

手写Triton Kernel跑通Qwen3.5-0.8B:21算子拆解与CUDA Graph优化 说实在的当我决定做“纯手写 Triton kernel 跑通 Qwen3.5-0.8B”这个项目时身边不少人觉得我是在给自己找麻烦。毕竟现在 vLLM、TensorRT-LLM 这些推理框架已经够成熟flash-attention 一个接口就能把 attention 算得又快又省显存何苦从算子层面一个个手写但我的想法很简单如果只停留在“调库”的层面你永远不知道 GPU 上到底发生了什么。这篇是系列第一篇先把整个项目的地基打好——包括 21 个算子的拆解、split-K 的切分思路以及 CUDA Graph 的接入方式。目标读者是那些对大模型推理原理感兴趣、想自己动手折腾 kernel 的工程师和研究者哪怕你现在刚接触 Triton也能从这篇里找到一条相对清晰的路径。这个项目的核心不是追求“跑得比 vLLM 快”而是追求“每个算子都在自己掌控之中”。0.8B 参数的模型规模不大不小正好适合用来做全量手写实验它既有完整的 transformer 结构又不会像 7B 模型那样让编译和调试周期长到令人崩溃。我最终把整个推理过程拆成了 21 个算子并在这 21 个算子的基础上叠加了 split-K 和 CUDA Graph 两层优化。下面我会按照实际推进的顺序把设计思路、算子清单、关键实现细节和踩过的坑逐一讲清楚。1. 项目全景为什么选择手写 Triton kernel1.1 一个 0.8B 模型的推理里究竟藏了多少算子很多人对“算子”这个词有误解觉得它是个很高端的 GPU 专属概念。其实算子就是“对一个数据张量做一次特定计算的单元”这个词也不只是深度学习领域在用——机器视觉里的 Halcon 算子、图像处理里的 Sobel 算子、拉普拉斯算子本质都是同一个意思一个可复用的计算步骤。大模型推理里的算子就是把矩阵乘、归一化、激活函数、位置编码这些步骤逐一实现成能在 GPU 上高效运行的内核函数。Qwen3.5-0.8B 的标准结构是 embedding 层 若干层 transformer block每层包含 attention 和 MLP 最终归一化 lm_head。听起来不算复杂但真正落到 GPU 上执行时你会发现自己面对的是一个又一个细碎的计算步骤。哪怕一个 RMSNorm 都有“求均值 → 求方差 → 归一化 → 缩放”好几个阶段而 attention 内部的 QKV 投影、旋转位置编码、因果 mask、softmax、加权求和更是每个都不能少。我按“一次 decode 生成一个 token 的完整前向路径”来统计把可以融合的算子合并后最终得到 21 个必须单独实现或至少单独考虑的算子。这 21 个算子覆盖了三个层面张量搬移类embedding 查表、KV cache 读写、计算密集类线性层、attention、MLP、规约与采样类softmax、topk。每一类在 Triton 里的实现难度和优化空间完全不同。张量搬移类看着简单但访存模式没写好会让整个 kernel 变成显存带宽瓶颈计算密集类需要你认真设计分块大小让 tensor core 真正转起来规约类则要处理好数值稳定性否则 fp16 下一不小心精度就崩了。1.2 手写 kernel 的价值与代价既然 flash-attention 现成的那么香为什么还要手写我的理由有两条。第一只有手写过一遍你才能真正理解那些第三方库在背后做了什么。flash-attention 的 online softmax、分块策略、重计算技巧如果只是当黑盒用遇到性能问题或者精度问题你根本不知道从哪排查。第二手写 kernel 给了你最大的定制自由度。比如 KV cache 怎么布局、mask 怎么处理、量化怎么融进算子这些在现成框架里往往只能接受框架的设定而自己写 kernel 时你可以在任意一层做想做的优化。但代价也真实存在。Triton 相比 CUDA 确实降低了不少心智负担你不需要手动管理线程块和共享内存的细节但调度策略、分块大小、数值精度这些仍然要自己操心。这个 0.8B 模型我完整调通花了大约两周其中一半时间都耗在“结果不对但不知道是哪个算子出了问题”的调试阶段。而且 21 个算子的编译和调优是个反复迭代的过程每改一次 kernel 参数都要重新编译、跑测、对比。所以如果你只是想快速把一个小模型跑起来做业务手写 kernel 绝对不明智但如果你想真正吃透推理性能的来龙去脉这笔时间投入非常值。2. 21 个算子的拆解与 Kernel 设计2.1 从模型结构出发把算子清单列清楚我强烈建议任何做同类项目的人第一步不要急着写代码而是先把模型结构的计算图拆成一张算子表格。这件事做扎实了后面写 kernel 才有据可依。以 Qwen3.5-0.8B 为基准我整理的 21 个算子如下表所示序号算子名称计算类别功能说明1Embedding访存词表索引查表取 token 对应的高维向量2RMSNorm_Attn归一化Attention 块前的 Root Mean Square 归一化3QKV_Proj计算密集输入向量经矩阵乘产生 Q、K、V 三组向量4RotaryEmbedding逐元素对 Q、K 施加旋转位置编码5KV_Cache_Update访存将当前 token 的 K、V 写入缓存6KV_Cache_Lookup访存读取历史 K、V 缓存用于注意力计算7QK_Matmul计算密集Q 与 K 矩阵乘得到注意力分数8Causal_Mask逐元素施加因果掩码屏蔽未来位置9Softmax规约注意力分数归一化为概率分布10PV_Matmul计算密集注意力概率与 V 矩阵乘得到上下文向量11Attn_Output_Proj计算密集上下文向量通过输出投影矩阵12Residual_Add_1逐元素Attention 输出与输入做残差相加13RMSNorm_MLP归一化MLP 块前的归一化14Gate_Proj计算密集门控分支线性投影15Up_Proj计算密集上投影分支线性投影16SiLU_Activate逐元素对门控分支施加 SiLU 激活函数17Down_Proj计算密集合并两条分支的线性投影18Residual_Add_2逐元素MLP 输出与输入做残差相加19Final_RMSNorm归一化最终输出层前的归一化20LM_Head计算密集映射到词表大小的 logits21TopK_Sample规约采样从 logits 采样选择下一个 token看到这张表你应该能感觉到所谓“跑通模型”其实就是把这 21 个算子串成一条流水线。其中序号 3 的 QKV_Proj 和序号 7、10 的两个 Matmul三者贡献了绝大部分计算量而序号 2、4、8、16 这类逐元素或小规约算子计算量不大但数量多恰恰是 kernel 启动开销的大头。这也是后文要引入 CUDA Graph 的直接原因——把 21 次启动合并成一次图执行。2.2 三类算子的 Triton 实现策略面对这 21 个算子我在实现时按计算特征分了三条线。第一条线是 Embedding 和 KV cache 相关操作重点是解决访存效率。以 Embedding 为例输入是一组 batch 的 token id做的事情就是拿 id 去 token embedding 表里按行取向量。最简单的方式是写一个逐行拷贝的 kernel但更好的做法是让相邻线程访问相邻显存充分利用 cache line。Triton 里我直接用tl.load配合一个指针偏移数组实现批量查表注意 embedding 表的行数可能很大block 大小不要设太小否则访存利用率上不去。第二条线是计算密集的矩阵乘类算子。QKV_Proj、Gate_Proj、Up_Proj、Down_Proj、LM_Head 本质上都是 matmul。Triton 里调用tl.dot就能让 tensor core 干活但要想性能好关键在于 BLOCK_M、BLOCK_N、BLOCK_K 这三个参数的组合。我实测下来0.8B 模型里这几个矩阵乘的维度不算大BLOCK_M 取 32 或 64 就够了BLOCK_N 跟输出维度相关通常 64 或 128BLOCK_K 取 32 或 64。太大反而会因为矩阵太小导致尾部浪费严重。另外fp16 输入下tl.dot默认用 fp32 累加这一点非常重要直接避免了一部分精度损失。第三条线是 Softmax、RMSNorm、SiLU 这类逐元素或规约算子。它们逻辑简单真正的坑在于数值稳定性和融合策略。Softmax 要分两步先求行的最大值再算 exp 和归一化中间不能省掉减去最大值这一步否则 fp16 下极容易出现 NaN。RMSNorm 没有 softmax 那么娇气但平方和累加时也要用 fp32 中间变量。SiLU 激活函数本质是x * sigmoid(x)可以直接写成一个逐元素 kernel更好的做法是把它融合到 Gate_Proj 的输出处理环节省一次显存读写。2.3 容易踩坑的“小算子”处理别看 Embedding、Residual_Add、Causal_Mask 这些算子代码量不大实际踩坑最多的反而是它们。第一个坑是数据类型的隐式转换。Triton kernel 里如果输入是 fp16中间计算默认会保持 fp16但残差相加这个操作我建议先转成 fp32 算完再加回去否则多轮残差积累下来精度损耗会慢慢变大最终生成结果可能和 PyTorch 参考实现对不上。第二个坑是 mask 的实现方式。Causal_Mask 不能真的去创建一个(seq_len, seq_len)的 mask 张量那太浪费显存了。正确做法是在 QK_Matmul 之后、Softmax 之前用tl.where基于位置索引直接生成掩码。Triton 里可以用offs_q[:, None] offs_k[None, :]这种广播比较来生成布尔掩码然后对 mask 位置填一个负大数比如 -1e38。这里有个细节负大数到底填多少取决于你后面是 fp32 还是 fp16 的 softmaxfp16 下填 -1e4 都可能溢出我建议直接填 -3e4 以上且小于 fp16 能表示的最小负数的量级实测 -1e38 在大多数情况下安全。第三个坑是 KV cache 的 layout。KV cache 的维度一般是(num_layers, batch, num_kv_heads, seq_len, head_dim)但如果用 Triton 每次单独读写一个 token 的 K、V你会发现访存效率很低。我在实际操作中做了一个调整把 KV cache 里同一个 head 的 seq 维度连续存放这样 KV_Cache_Update 写入时和 KV_Cache_Lookup 读取时都能用连续的内存访问模式整个 attention 计算的访存开销明显下降。这种底层布局的调整恰恰是手写 kernel 才有的自由度。3. split-K把大矩阵乘法切出并行度3.1 什么时候需要 split-K先澄清一个概念split-K 指的是在矩阵乘的 K 维即两个矩阵相乘时内部累加的那个维度也叫规约维上做切分让不同的计算单元分别计算部分累加和最后再把部分和加起来。常规的矩阵乘分块是在 M 和 N 维度上分块K 维度由单个线程块内部串行累加。当 M、N 维度都比较小、而 K 特别大时单靠 M×N 的分块没法用满 GPU 的并行度这时候 split-K 就能把 K 维度切成多个块让更多线程块参与计算。具体到 Qwen3.5-0.8B 这个模型最典型的场景是 decode 阶段。此时 QK_Matmul 的形状是(1, num_heads * head_dim)乘(num_heads * head_dim, seq_len)M 维度只有 1N 维度是 seq_len并行度天然受限。如果 seq_len 还不算太长比如只有几十或几百那 GPU 的 SM 大多数都在空转。另一个场景是 LM_Head它的形状是(1, hidden_dim)乘(hidden_dim, vocab_size)M 同样为 1vocab_size 虽然能到几万但如果 batch 很小整体并行度依然不足。这两个地方是我实际应用 split-K 最多的地方。3.2 实现思路与调参在 Triton 里实现 split-K我不会真的去手动拆 K 然后管理多个 block——Triton 的编程模型天然支持你通过tl.dot的块内遍历来做但要实现 split-K 的核心还是把 K 维分给多个 program。一个典型的写法是把PID_K作为 K 维的块索引每个 program 负责 K 维的一个切片计算得到局部累加结果后再通过 atomic add 或单独的第二阶段 kernel 做累加。我在项目里用的方案是 atomic add。Triton 里可以用tl.atomic_add把局部结果累加到输出矩阵上但要注意初始化输出矩阵为零以及保证浮点累加的确定性。这里有个细节如果直接tl.atomic_add多个局部结果即使每个局部结果内部顺序一致多个 program 之间到达的顺序是不确定的所以最终结果和单次累加的数值会有细微差别。对 fp16 来说这个差别通常在可接受范围内但如果你需要完全确定性的结果就需要把局部结果写到一个中间张量然后让一个单独的 kernel 做规约累加。为了性能我选了 atomic add为了稳定性我额外写了一个小的规约 kernel 做分期累加实际跑下来两者差异不大但精度确定性好了很多。split-K 的切分数目不是越多越好。我实测发现切分数目超过 8 之后atomic add 的竞争开销开始盖过并行收益。以 QK_Matmul 为例K 维度实际是 head_dim通常是 128这种情况下 split-K 切成 2 到 4 份就够了而 LM_Head 的 K 维度是 hidden_dim可能到 2000 多切成 4 到 8 份比较合适。你在自己的项目里需要根据矩阵实际形状和 GPU 核心数做一组消融实验重点观察不同 split 数下的耗时曲线通常会有一个明显的拐点。3.3 split-K 的代价与取舍split-K 从来不是免费的。它最直接的代价是额外的原子加操作或第二阶段的规约 kernel这部分开销会随 split 数量的增加而上升。另一个代价是 L2 cache 的压力——当多个 program 同时读取同一个输入矩阵的不同 K 切片时如果输入矩阵不在 cache 里访存带宽就成了瓶颈。所以 split-K 只在“必要的并行度不足”时才引入M、N 维度本身已经足够大的场景下强行用 split-K反而会让性能变差。举个例子在 prefill 阶段输入序列可能已经很长attention 矩阵乘的 N 维度是完整的 seq_len并行度并不匮乏这时候我就完全不用 split-K。而 decode 阶段序列很短、batch 通常是 1就是 split-K 发挥价值的主战场。这个取舍你需要在自己的机器上用 nsys 或简单的 torch profiler 验证不要盲信任何一个固定配置。另外还要提一句如果用了 CUDA Graph后面会讲split-K 产生的额外 kernel 数量也会被图捕获所以它对启动开销的影响几乎可以忽略这也是把两者放在一起做优化的原因之一。4. CUDA Graph把 21 次启动压成 1 次4.1 为什么需要 CUDA Graph逐 token 生成是 LLM 推理的基本工作方式。每次 decode 都要从 CPU 侧发一个 kernel 到 GPU而这个 kernel 通常非常小执行时间可能只有几微秒到几十微秒。如果是手写算子21 个算子串行启动 21 次每次启动的 CPU 开销、内核调度开销、参数传递开销全部叠加起来可能比 kernel 实际执行时间还长。我最初跑通 21 个算子串行版本时单次 decode 的 GPU 计算时间只有 1 毫秒左右但加上 21 次启动的 CPU 开销整体耗时接近 4 毫秒这就是典型的小 kernel 启动开销主导的场景。CUDA Graph 解决的就是这个问题。它的思路是把一串 kernel 的启动过程提前“录制”下来之后每次执行只需要一次图启动CPU 侧的调度开销大幅降低。用大白话说常规方式是 CPU 逐条喊口令、GPU 逐条执行CUDA Graph 则是把一整本口令集提前录好之后 CPU 只需要喊一声“开始”GPU 就自动跑完一整段。这在数学上没改变计算本身但省掉了大量指令往返的延迟。4.2 捕获流程与 Triton 的配合Triton kernel 编译出来之后本质上也是普通 CUDA kernel所以完全可以放进 CUDA Graph 里捕获。捕获的基本流程是先准备一个独立的 CUDA stream然后在这个 stream 上执行一遍完整的推理过程同时让 CUDA Graph 记录所有 kernel 启动之后把这个图保存下来重复使用。在 PyTorch 里torch.cuda.graph上下文管理器帮我们封装好了这套流程用起来相当方便。我在项目里接入时的实际代码如下简化版import torch from torch.cuda import graph as cuda_graph # 提前分配好所有中间张量和输出张量 static_inputs {...} # 提前 fixed shape 的输入 g cuda_graph.CUDAGraph() # 预热确保所有 Triton kernel 已完成编译 for _ in range(10): run_inference(static_inputs) # 开始捕获 with torch.cuda.graph(g): run_inference(static_inputs) # 之后每次 decode 只需要 replay g.replay()这里面最关键的前提是捕获阶段使用的所有张量地址必须是“静态”的也就是不能动态分配新的显存。为了做到这一点需要在捕获之前就固定好所有中间缓冲区的大小和地址。我在实际项目中给每个算子的输出都预分配了一块显存比如 attention 输出的上下文向量、MLP 中间激活等这样图捕获之后每次 replay 不会触发新的内存分配也就不会导致图失效。这个“静态化”是 CUDA Graph 接入时最繁琐但最必要的工作。4.3 动态 shape 问题的解决CUDA Graph 最大的限制在于它捕获的是固定的 kernel 启动序列和固定的显存地址一旦 shape 变了整个图就必须重新捕获。而 LLM 推理里 shape 经常变——decode 阶段 sequence length 会随着生成而增长batch 大小也可能动态变化。这个问题不做处理CUDA Graph 根本没法稳定复现。我的处理方式是把模型推理的主要路径设计成“shape 不变”的。具体来说我给 KV cache 分配最大长度的缓存每次 decode 时把当前 seq_len 作为固定值传入 kernelkernel 内部用这个值限制实际参与计算的范围而所有中间张量的 shape 在捕获时就是最大值。这样虽然计算上会多算一些被 mask 掉的无效位置但换来的是图可以一直复用不需要反复捕获。对 0.8B 模型来说这种“算一点多余位置”的代价远远小于反复捕获图带来的开销。另一个细节是 batch size 动态变化的处理。我固定 batch size 为 1这也是逐 token 生成最常用的模式这样整个图的输入输出实际上只需要两个指针一个输入 token id一个输出 logits。如果你确实需要动态 batch一种方案是预分配最大 batch 的显存然后用 padding 的方式把不足 batch 的位置用 mask 处理掉思路和 sequence length 的处理一模一样。4.4 实测效果接入 CUDA Graph 之后我在 Qwen3.5-0.8B 上做了几组对比实验场景是 batch size 为 1 的逐 token 生成模型精度使用 fp16。单次 decode 的耗时从手写串行版本的约 4 毫秒降到了约 1.2 毫秒其中 GPU kernel 本身的执行时间和之前实际上差不多省掉的 2.8 毫秒几乎都是 CPU 启动开销。换句话说CUDA Graph 不会让你的 kernel 变快但它能让你的“整体系统”变快因为硬件不再需要反复等待指令。如果你想让收益更明显还可以在 CUDA Graph 外面配合一个更高效的前端。比如模型生成时连续调用graph.replay()就能省掉 Python 层的多次张量搬运。我实际测试过在同一个 GPU 上不依赖任何推理框架单张 24GB 显存可以轻松跑起来这个 0.8B 模型单 token 生成延迟在毫秒级。这个成绩当然没法跟 vLLM 的连续批处理比吞吐但在“从零手写”的场景下已经足够证明方案的可行性。5. 常见问题与排查技巧实录5.1 Triton kernel 结果不对怎么定位是哪个算子出错这是手写 kernel 项目里最痛苦的问题。21 个算子串联执行中间任何一步出错最终产出的 logits 都会偏。我的排查方法分三步走。第一步准备一个 PyTorch fp32 的参考实现每个算子都用最朴素的矩阵运算或 PyTorch 原生算子计算得到一份“标准答案”。第二步把每个 Triton kernel 单独拎出来用相同的输入张量同时跑 Triton 版本和 PyTorch 参考版本逐算子对比输出。第三步用torch.allclose设置合理的 atol 和 rtol 做自动对比不要用肉眼去看那些张量数字。具体操作时我会在每个算子的 Triton 实现里加上一个开关比如环境变量DEBUG_ON开关打开时 kernel 输出直接与参考实现做差并把差值的绝对均值和最大值打印出来。这样一旦最终结果不对就能快速定位是哪个算子出的偏差。在 Qwen3.5-0.8B 上我实际踩过几个精度坑第一个是 fp16 下tl.dot的输出累加需要用 fp32否则几次矩阵乘之后误差会累积到生成结果完全不可用第二个是 softmax 的 exp 在输入较大时直接溢出必须做 max 平移第三个是残差相加的中间变量精度前面提过也要转 fp32。5.2 CUDA Graph 捕获失败的常见原因CUDA Graph 捕获失败是很常见的问题尤其是手写 kernel 项目里。最典型的原因是捕获过程中发生了动态内存分配。比如某个 kernel 的输出张量是在推理函数内部通过torch.empty创建的捕获时 PyTorch 就会向 CUDA 分配器申请新的显存块这会导致捕获直接失败或产生性能警告。解决方式就是我前面提到的静态化所有中间张量都在捕获前分配好推理函数只读取和写入这些预分配好的 buffer。第二个高发原因是 kernel 内部有设备同步操作比如torch.cuda.synchronize()、tensor.item()、tensor.cpu()这类操作它们会让 CPU 等待 GPU破坏 CUDA Graph 的录制逻辑。我在调试的时候也踩过某个算子因为想打日志而在内部调用了item()捕获时直接报错。排查这类问题有个技巧——捕获时报错信息里如果出现 “operation not permitted during capture” 之类的字样基本就是有同步或动态内存操作顺着调用栈把对应代码静态化即可。第三个原因是显存不足。捕获时 CUDA Graph 需要把 kernel 的启动信息记录下来这部分信息本身也要占显存虽然不大但如果在捕获前显存已经接近上限也会失败。这个在 0.8B 模型上基本不会碰到但如果你后续扩展到更大模型就需要留意。5.3 性能瓶颈定位与调优建议跑通只是第一步跑得快才是目标。我建议你在做完首版之后立刻用 profiling 工具看一遍时间分布。最简单的方式是用 PyTorch 自带的 profiler或者用更专业的 Nsight Systems。打开时间线之后你会清楚看到每个 kernel 的耗时、kernel 之间的间隙、以及 CPU 在每次 launch 上的开销。我在首版就发现 QK_Matmul 和 MLP 的几个小矩阵乘时间占比最高这才下定决心引入 split-K。调优的时候我一般按这个顺序来先看有没有可以融合的逐元素算子比如残差相加和下一个 RMSNorm 的输入加载可以合并再看矩阵乘的分块参数是否合理用 Triton 自带的triton.testing.do_bench对不同 BLOCK 大小做一组快速实验最后才考虑 split-K 这类更复杂的改动。记得每做一步优化都要回归精测一次防止性能上升但精度下降的隐性风险。还有一个小技巧Triton 的编译缓存有时会导致你改了参数却不生效。我在迭代调参时遇到过明明改了 BLOCK_SIZE 但耗时不变的情况最后发现是编译缓存命中了旧的二进制。遇到这种情况删掉~/.triton/cache目录再跑就好。如果你希望每个 kernel 的编译过程更透明可以在首行设置TRITON_PRINT_AUTOTUNING1环境变量这样 Triton 会输出每个 kernel 的调优结果。另外关于 split-K 的调参我总结了一个简单的经验先用 profiler 看某个矩阵乘的 occupancy如果计算核心利用率低于 50%就有引入 split-K 的空间如果已经高于 70%那 split-K 带来的收益会很有限甚至会因为原子加开销倒挂。这个判断方法对我很有用推荐你也试试。6. 系列后续与扩展思路这篇先把 21 个算子的骨架、split-K 和 CUDA Graph 讲完了但对一个完整可用的推理系统来说这只是第一步。后续我计划系统聊几个方向一是量化0.8B 模型可以用 INT8 或 INT4 大幅降低显存占用这需要把大部分算子都改成支持定点计算版本二是连续批处理现在的实现是单 batch 逐 token 生成加上连续批处理才能真正和主流推理框架比吞吐三是 long context 的优化当前 KV cache 是朴素实现seq 变长后注意力计算的开销会显著上升后续可以考虑类似 PagedAttention 的分页管理思路。手写 kernel 这条路上最难的不是某一个算子的实现而是把 21 个算子组合成一条高效、稳定、可维护的流水线。很多细节需要在真实跑测中一点点摸出来我分享的这些也只是一家之言。如果你也在折腾类似的项目欢迎在评论区交流你的 split-K 参数选择或者 CUDA Graph 捕获时遇到的奇葩问题我踩过坑的地方或许能帮你少走几步弯路。
返回列表