尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

Triton手写21个算子:从PyTorch到CUDA Graph的推理优化实战

发布时间:2026/9/9 13:52:21

资讯中心
01
ARTICLE

Triton手写21个算子:从PyTorch到CUDA Graph的推理优化实战

Triton手写21个算子:从PyTorch到CUDA Graph的推理优化实战
1. 为什么有这篇文章从“能跑”到“跑得快”做推理优化的朋友应该都有同感用 PyTorch 把模型跑通只是第一步真正磨人的是把 kernel 一层层拆开、重写、压榨到带宽和算力都用满。我这次选的靶子是 Qwen3.5-0.8B一个 8 亿参数左右的小模型。选它的原因很直接单卡能放下推理延迟对 kernel 实现的敏感度又足够高任何一处偷懒都会在耗时上暴露出来特别适合用来验证手写 Triton kernel 的整套方法论。这篇文章是系列的第一篇核心做三件事把模型推理路径上真正需要的 21 个算子逐个用 Triton 重写把 split-K 这种归约优化的思路讲透最后用 CUDA Graph 把整个解码过程的 CPU 启动开销压到极限。目标读者是已经会写简单 Triton kernel、想进一步把一个小模型从底层吃透的工程师。如果你只是调包跑模型这篇文章可能过于底层但如果你想理解推理引擎里那些算子到底在做什么这里面每一段都值得读两遍。先说结论免得后面看得没耐心21 个算子全部手写并融合后在不牺牲精度的前提下单 batch 解码延迟比原生 PyTorch eager 模式快约 2.6 倍比 torch.compile 快约 1.5 倍核心收益来自三块——算子融合减掉了大量中间张量的读写、split-K 把注意力分数和 layer norm 这类归约计算打满了 SM、CUDA Graph 把每 token 的 kernel 启动开销从百微秒级降到了个位数微秒级。1.1 手写 Triton kernel 的目标与边界在开工之前我先给自己划了三条边界。第一不重写全部计算只重写推理路径上被反复调用、且 PyTorch 实现有明显性能浪费的算子第二所有 kernel 必须是纯 Triton 实现不用 cuBLAS 的隐含调用避免“手写”名不副实第三数值精度不能比 PyTorch 的默认实现差太多Kernel 内部允许用 FP32 累加输入输出保持 FP16/BF16。这三条边界很重要。一个常见误区是“手写 kernel 就要把 100 个算子全部重写”但实际推理时像权重加载、量化参数预处理这些一次性操作根本不需要进 kernel。真正需要优化的是每条 token 路径上反复执行的算子。另一个边界是精度。我在试第一版的时候发现如果 RoPE 里直接用 FP16 算三角函数位置一长误差就会累积表现为困惑度轻微上升。后来统一改成“输入 FP16中间用 FP32 计算输出再转回 FP16”问题就消失了。这里也提醒各位Triton 里tl.cos、tl.sin这些函数目前还是 FP32 实现你在代码里手动转换一下就行千万不要为了省几条指令让精度白白损失。1.2 0.8B 模型推理路径上的算子分布先盘一下 Qwen3.5-0.8B 这个规模的模型推理路径上有哪些计算量占比高的算子。按层来看Transformer block 内部主要是四个部分输入 LayerNorm、Self-Attention含 QKV 投影、RoPE、注意力分数、Softmax、输出投影、残差连接、Feed-Forward 网络SwiGLU 结构含 gate 和 up 两个分支。再加上最开始的 token embedding最后的 LM Head 和采样逻辑。如果完全不融合PyTorch 默认会把这些拆成 40 到 50 个小算子。我这次通过算子融合把重复的访存去掉压缩到了 21 个。核心思路是RMSNorm 和 RoPE 尽量融合进前面的线性层或注意力计算里能少读一遍中间张量就少读一遍Softmax 和 attention score 的计算尽量留在同一份 SRAM 数据里完成。下面这张算子清单就是 21 个算子的完整划分后面几个小节分别讲实现。这张表我建议先存一下后面看代码的时候对照着来。2. 21 个算子LLM 推理的最小 kernel 集合2.1 完整的算子清单与职责这 21 个算子不是随意拆的而是按“计算类型 融合边界”两个维度来划分的。我先列出全部算子再解释为什么这么切。编号算子名融合后职责对应的原始操作1embedding_lookupToken ID 查表nn.Embedding2qkv_projQ/K/V 三个线性层合并3 个 Linear3q_ropeQ 的旋转位置编码RoPE4k_ropeK 的旋转位置编码RoPE5qk_scoreQK^T 点积torch.matmul6softmax_maskMask Softmaxmasked_fill softmax7pv_aggPV 加权求和torch.matmul8attn_out_projAttention 输出投影Linear9attn_residual残差连接Addx attn_out10mlp_gate_upSwiGLU 的 gate 和 up 两个分支合并2 个 Linear11swish_multSwish 激活与逐元素乘SiLU mul12mlp_down_projdown 投影Linear13mlp_residual残差连接Addx mlp_out14ln1Attention 前 RMSNormRMSNorm15ln2MLP 前 RMSNormRMSNorm16logits_headLM Head反 embeddingLinear17argmax_sample采样取最大概率 tokentorch.argmax18cache_kv_writeKV Cache 写入索引赋值19cache_kv_readKV Cache 读取索引读取20causal_mask_gen因果掩码生成预计算 mask21scale_softmaxAttention 缩放与归一化scale softmax你可能会问为什么 qkv_proj 是融合的而 q_rope 和 k_rope 要分两个算子原因是 Q 和 K 的变换逻辑不同Q 的旋转是每一对维度旋转K 的旋转还涉及 KV Cache 的读取时机。在 decode 阶段Q 只有当前 token 一个K 要追加到缓存里所以两者融合边界天然不同拆开写反而清晰。2.2 算子之间的融合决策少读一遍中间张量融合的核心逻辑只有一个减少全局内存的往返次数。GPU 计算密度高但内存带宽是稀缺资源省一次张量读写省下的时间往往比省几次计算还多。以 qkv_proj 为例。PyTorch 里如果拆成三个 Linear那么 Q、K、V 三个中间结果都要先写回全局内存再分别被后面的 RoPE 和注意力读取。融合后一次 kernel 同时计算三个投影Q、K、V 都留在寄存器或 SRAM 里直接消费。实测数据是单纯这一个融合在 0.8B 模型上 decode 延迟少了约 12%。又比如 softmax_mask。PyTorch 传统写法是out x.masked_fill(mask 0, float(-inf)); out out.softmax(dim-1)这里涉及两次全局内存的读写。Triton 版本里mask 和 softmax 在一个 kernel 内完成tl.where判断后直接在寄存器里做指数减最大值、归一化中间结果不落内存。还有一组容易被忽略的融合是ln1和qkv_proj。因为 Attention 前的 RMSNorm 结果立刻被 QKV 投影使用RMSNorm 完全没必要把归一化后的张量写回全局内存。可以直接把 RMSNorm 的 kernel 与 qkv_proj 合并只输出 Q、K、V。这个融合我看很多代码库都没做可能是为了保持代码结构清晰但实测能省 3% 到 5% 的时间在延迟敏感的推理场景里不算小数字。2.3 手写 RMSNorm 和 RoPE 时的工程细节先看 RMSNorm它比 LayerNorm 简单一些不计算均值只做方差归一化。但工程实现里有两个坑。第一个坑是统计量计算与归一化拆成两个 kernel 会导致多读一遍输入第二个坑是 FP16 累加误差。我的做法是在一个 kernel 里同时算出均方根并归一化累加过程用 FP32。Triton 代码大致是这样torch.no_grad() def rms_norm_kernel(X, W, Y, x_row_stride, w_row_stride, y_row_stride, n_cols, eps, BLOCK_N: tl.constexpr): row_idx tl.program_id(0) cols tl.arange(0, BLOCK_N) x_ptrs X row_idx * x_row_stride cols w_ptrs W cols x tl.load(x_ptrs, maskcols n_cols, other0.0).to(tl.float32) w tl.load(w_ptrs, maskcols n_cols, other0.0) var tl.sum(x * x, axis0) / n_cols rstd 1.0 / tl.sqrt(var eps) y x * rstd * w y_ptrs Y row_idx * y_row_stride cols tl.store(y_ptrs, y.to(X.dtype.element_ty), maskcols n_cols)关键点是把x先转到 FP32tl.sum的累加默认也是 FP32最后存的时候再转回 FP16这样精度不会丢。RoPE 的实现要特别注意“频率因子”的生成。通常有两种方式提前生成一个cos_cache/sin_cache传到 kernel 里或者在 kernel 内部用公式现算。我的建议是提前生成缓存。因为 decode 阶段每个 token 调用 RoPE 的时候位置是固定的反复计算三角函数完全没必要。Triton kernel 里只需要按列的奇偶位置读取 cos、sin 缓存做旋转。torch.no_grad() def rope_kernel(Q, K, Q_out, K_out, cos, sin, q_batch_stride, q_head_stride, q_dim_stride, head_dim, pos, BLOCK_D: tl.constexpr): pid tl.program_id(0) # flatten over (batch * num_heads * seq) ... half_d head_dim // 2 d_idx tl.arange(0, BLOCK_D) within_head d_idx head_dim # 前半部分存偶数维度后半部分存奇数维度 even_idx d_idx // 2 cos_vals tl.load(cos pos * head_dim even_idx, maskwithin_head, other1.0) sin_vals tl.load(sin pos * head_dim even_idx, maskwithin_head, other0.0) # 判断当前维度是偶数还是奇数 is_even (d_idx % 2) 0 partner_idx tl.where(is_even, d_idx 1, d_idx - 1) ...这里“partner”的概念容易写错旋转位置编码里偶数维度和它下一个奇数维度组成一对旋转是二维平面上的旋转不是逐元素乘。我第一次写的时候按逐元素乘处理loss 直接崩了后来画了个二维旋转变换图才反应过来。建议大家在代码里保留is_even这个判断脑子不清楚的时候看一眼逻辑。3. split-K让归约类算子吃满 GPU3.1 split-K 到底解决了什么问题GPU 上有大量“归约类”算子经典代表是矩阵乘法中对 K 维度的累加、LayerNorm/RMSNorm 中对输入向量求和、Attention score 的点积、Softmax 中的最大值与指数和。这些操作的核心特征是计算量不大但涉及的数据要么从一个维度反复读取要么需要跨线程协作归约。如果在 Triton 里不做任何处理一个简单的 kernel 会用一个 program 处理一行或一个矩阵块归约发生在单个 program 内部。问题来了当矩阵的 M 维行数较小时比如 decode 阶段 M1你只启动了几十个 program远填不满 GPU 上的上百个 SMSM 大部分时间空转。split-K 的思路是把 K 维度主动切成 K 份每一份由一个独立的 program 做部分归约最后再做一次小的归约汇总把并行度提上去。举个例子decode 阶段 QK^T 的形状是[1, num_heads, 1, head_dim] [1, num_heads, head_dim, seq_len]理论计算量很小但如果不做 split-K一个 head 的计算只分给一个 program整个 GPU 可能只用了 10% 的算力。把 K 维度 split 成 8 份后一个 head 有 8 个 program 同时干活SM 占用率自然上去。3.2 Triton 里 split-K 的常规写法Triton 给了我们一个非常方便的归约原语tl.atomic_add配合tl.atomic_max可以实现跨 program 的轻量协作。以 attention score 的 QK^T 为例我写了一个 split-K 版本的 kerneltorch.no_grad() def qk_score_split_kernel(Q, K, Score, q_ptr, k_ptr, score_ptr, num_heads, head_dim, seq_len, scale, SPLIT_K: tl.constexpr, BLOCK_K: tl.constexpr): pid_m tl.program_id(0) # batch * head pid_k tl.program_id(1) # split index # 每个 program处理 head_dim 里的 BLOCK_K 段 block_start pid_k * BLOCK_K offsets block_start tl.arange(0, BLOCK_K) # 这里M1直接处理一行 q_vals tl.load(Q pid_m * head_dim offsets, maskoffsets head_dim, other0.0).to(tl.float32) k_vals tl.load(K pid_m * head_dim offsets, maskoffsets head_dim, other0.0).to(tl.float32) part_sum tl.sum(q_vals * k_vals, axis0) * scale # 注K维度切分后每个program只算部分和需要原子加到全局地址 tl.atomic_add(Score pid_m, part_sum)当然这只是 QK^T 一处的 split-K。在完整注意力实现里我会把 QK^T 与 softmax 融合到同一个 kernel 中softmax 本身也可以做 split-K但需要保存部分最大值和部分指数和写起来更复杂。这里留个作业如果你的 seq_len 很长可以自己试着把 online softmax 的分块版本实现一下原理是每个 split 存(max, sum_exp)最后做一次 merge。3.3 split-K 的关键参数选择split-K 不是越多越好。启动的 program 数量等于num_heads * SPLIT_K当这个乘积远大于 SM 数量时会带来额外的调度开销和原子操作冲突。以 A100 为例SM 数是 1080.8B 模型一般是 16 到 32 个 head所以 SPLIT_K 取 2 到 4 通常就够。我实测下来decode 阶段 head_dim128 时SPLIT_K4 比 SPLIT_K8 快了大约 5%因为 8 份引起的原子竞争比收益还大。还有一个小细节split-K 需要保证最终结果的确定性。原子加法的顺序不确定对精度敏感的场景可能造成逐次运行结果略有差异。我的解决方案是对 score 做 split-K 原子操作时在 kernel 最后额外读回检查一遍如果与 CPU 参考实现差距大于某个阈值就报警。实际跑下来FP32 累加 原子加的误差在 1e-5 量级完全可接受。4. CUDA Graph把 CPU 参与降到最低4.1 为什么 decode 阶段需要 CUDA Graph在 GPU 推理里真正执行计算的耗时往往只占一部分CPU 启动 kernel 的开销同样可观。逐 token 生成时每个 token 要串行执行十几个到二十几个 kernel每个 kernel 启动都有固定的 CPU 开销约 3 到 10 微秒。一个 0.8B 模型每 token 生成大约需要执行 20 个 kernel如果都用原生 API 一层层启动光启动开销就是 60 到 200 微秒。这在早年的优化里还没那么致命但在现代 GPU 上一个 kernel 计算本身可能只要几十微秒启动开销占比就变得非常难看。CUDA Graph 的思路是把一串有依赖关系的 kernel 调用提前捕获成一个图之后每轮推理只需要重放这个图CPU 不再逐个启动 kernel而是把整个计算图一次性提交给 GPU。这在 GPT 这类生成场景里几乎是必做的优化。4.2 接入 CUDA Graph 的完整流程在 PyTorch 里接入 CUDA Graph 并不复杂但有几个容易踩的坑。标准的做法是三步捕获、重放、清理。第一步准备静态输入输出缓冲区。CUDA Graph 捕获到的 kernel 会绑定固定的内存地址所以参与计算的输入输出张量必须提前分配好且在整个捕获和重放周期内地址不能变。我是用torch.empty预分配了input_ids、kv_cache和logits每次解码只需要把新的 token 值拷贝进输入缓冲区而不是重新分配张量。第二步捕获。PyTorch 提供了一套torch.cuda.graph上下文管理器内部会自动开启捕获流。捕获时要把整个 decode 迭代的函数包进去包括 embedding、transformer层、logits、采样。这里有个容易犯的错误不要在捕获的代码里写任何 Python 层的 if/else 分支控制流因为图捕获只会记录第一次执行的路径。采样完成后根据结果决定是否终止这类逻辑要放到图外面。def decode_step(input_ids, kv_cache, ...): # 这里放所有 Triton kernel 的调用 logits model_decode(input_ids, kv_cache, ...) next_token sample(logits) return next_token # 预分配静态输入 static_input torch.empty((1, 1), dtypetorch.long, devicecuda) static_kv_cache torch.empty(...) g torch.cuda.CUDAGraph() # 预热确保所有 kernel 完成编译 decode_step(static_input, static_kv_cache, ...) torch.cuda.synchronize() with torch.cuda.graph(g): static_logits decode_step(static_input, static_kv_cache, ...)第三步重放。每次生成时先把新的 token 写入static_input然后调用g.replay()接着从static_logits里读结果。注意输入和输出的数据在 replay 之后才会更新同步点需要考虑好。static_input.copy_(token) g.replay() next_token postprocess(static_logits)4.3 捕获前后的常见坑第一个坑是捕获前的预热不充分。Triton kernel 有 JIT 编译过程第一次调用会触发 autotune 和编译如果不预热捕获到的图里会包含编译期的同步逻辑重放时会出错或性能极差。解决方法是先用真正的输入跑一遍确保所有 kernel 都已编译完成再开始捕获。第二个坑是动态 shape。CUDA Graph 是静态的一旦捕获块大小、grid 大小都固定了。如果模型支持变长输入需要按不同长度分别捕获多张图。对 0.8B 这种小模型我建议固定 seq_len 和 batch1按需要捕获几个不同长度的图即可。第三个坑是内存分配器与图的互操作。PyTorch 的缓存分配器可能在 replay 时给出不同的地址导致图失效。官方推荐捕获前调用torch.cuda.synchronize()并在捕获期间设置torch.cuda.memory._set_allocator_settings(expandable_segments:False)来确保稳定性。我自己在实际项目里固定 host 侧用copy_往里写数据没怎么遇到地址漂移但保险起见还是建议照做。第四个坑是“图内不要有返回值依赖”。如果你在 decode_step 里写了if next_token eos: break这个判断会在图捕获时被永久固化图重放时永远不会触发退出。正确做法是把判断放在 replay 之后由 Python 侧的循环控制是否继续。这也是 CUDA Graph 与 Python 控制流最大的区别。5. 实操过程与核心实现5.1 环境与依赖我的运行环境是单张 A100 40GBPyTorch 2.1Triton 2.2CUDA 12.1。如果你用的是别的卡比如 RTX 3090 或 4090基本也能跑只是 SM 数量和显存带宽不同会导致最优的 split-K 块大小变化需要重新 autotune。0.8B 模型本身很小显存主要被 KV Cache 和 activation 吃一些即使在 8GB 显存的卡上也问题不大。建议安装 Triton 时用官方 wheel而不是 PyTorch 自带的版本因为 Triton 迭代很快新版本对tl.atomic_add和tl.sum的优化有明显提升。我之前用 PyTorch 2.0 内置的 Triton 2.0 跑同一个 kernel比 2.2 慢了 15% 左右。5.2 全流程串联prefill 与 decode模型推理路径分成两个阶段prefill预填充和 decode解码。prefill 阶段处理用户输入的整段提示词并行度高计算量集中在矩阵乘decode 阶段一次只生成一个 token内存访问密集瓶颈在访存带宽和 kernel 启动开销。21 个算子在两个阶段的调用方式不同。prefill 阶段qkv_proj、mlp_gate_up 这些线性算子可以直接调用 Triton 的矩阵乘 kernel也适合把 seq_len 维度作为大 M 维decode 阶段则要切换为 M1 的 GEMV 或带 split-K 的专用 kernel。我的实现里维护了一张表标记每个算子在不同阶段使用哪个 kernel 变体方便后面做自动调度。prefill 阶段我还做了另一件事把 causal mask 提前生成好。因为 prefill 的 token 数量是已知的mask 可以预先算成一个 bool 矩阵放到全局内存kernel 里直接用避免在每层注意力里重复生成。这个 mask 在 decode 阶段不需要因为 decode 阶段只有一个新 token它应该看到之前所有 tokenmask 天然是全 1。5.3 21 个算子的 Triton 关键实现示例挑三个有代表性的算子细说。第一个是 qkv_proj。我把三个线性层合并成一个 kernel输入 x 的 shape 是[seq_len, hidden_size]权重分别是 W_Q、W_K、W_V输出 Q、K、V。Triton 里处理这种多输出的矩阵乘最简单的方式是每个 program 负责输出矩阵的一部分Q、K、V 各占一块同时计算torch.no_grad() def qkv_proj_kernel(X, W_Q, W_K, W_V, Q, K, V, seq_len, hidden_size, qkv_out_size, BLOCK_M: tl.constexpr, BLOCK_K: tl.constexpr): pid_m tl.program_id(0) pid_n tl.program_id(1) rm pid_m * BLOCK_M tl.arange(0, BLOCK_M) rn pid_n * BLOCK_N tl.arange(0, BLOCK_N) rk tl.arange(0, BLOCK_K) x_ptrs X rm[:, None] * hidden_size rk[None, :] x tl.load(x_ptrs) # 根据 pid_n 判断当前输出是 Q、K 还是 V 的分块 if rn[0] qkv_out_size: ...注意这里有个隐藏细节Q、K、V 的列维度是独立的所以当 pid_n 超出当前矩阵的列范围时要做 masking 或直接跳过。更好的做法是把 QKV 权重拼成一个[3 * head_num * head_dim, hidden_size]的大矩阵这样一次矩阵乘直接输出拼接结果再按视图切分。我最后采用的是后者代码更简洁而且因为 Triton 对矩阵乘有很好的 tile 优化把三个小矩阵合并成一个大矩阵乘吞吐更高。第二个是 pv_agg注意力输出聚合。Attention 权重 P 是在线的与 V 相乘后得到 attention 输出。这一步如果不做优化P 要写回全局内存V 要从 KV Cache 里读中间再加一次乘法访存往返很多。我在实现里把 qk_score、softmax_mask、pv_agg 三个算子合并成一个 flash-attention 风格的 kernel只保留最后一个输出写回全局内存。第三个是 cache_kv_write 和 cache_kv_read。这两个算子不值得写成复杂 kernel直接用 Triton 的指针偏移做即可。KV Cache 在 decode 阶段是关键路径读写的带宽利用率直接决定 decode 速度。我建议把 KV Cache 的布局从[batch, num_heads, seq_len, head_dim]改成[batch, num_heads, 2, seq_len, head_dim]或类似的结构让 K 和 V 分块时能按连续内存读写避免两个矩阵交叉导致缓存行利用率下降。5.4 使用 autotune 选择合适的 split-K 块大小Triton 提供tl.autotune装饰器可以自动搜索最优 kernel 配置。对 split-K kernel 来说配置项是 SPLIT_K 和 BLOCK_K 的组合。我第一次直接硬编码 SPLIT_K4在 4090 上跑的分数反而比 A100 差因为 4090 的 SM 数量和 L2 带宽不同最优参数不同。加上 autotune 之后每个 kernel 启动前会跑一次基准测试选最优配置。autotune 的坑在于它会增加启动前的搜索时间。在 decode 场景里如果每个 token 都重新 autotune那是灾难。正确做法是 warmup 阶段跑一遍 autotune之后把最优配置缓存下来实际推理时不再搜索。Triton 的 autotune 默认会对同一个 key通常是输入 shape做缓存所以只要你传的 batch size、seq_len 不变第二次调用时不会重新搜索开销很小。tl.autotune( configs[ tl.Config({SPLIT_K: 2, BLOCK_K: 64}, num_warps4), tl.Config({SPLIT_K: 4, BLOCK_K: 64}, num_warps4), tl.Config({SPLIT_K: 4, BLOCK_K: 128}, num_warps8), tl.Config({SPLIT_K: 8, BLOCK_K: 128}, num_warps8), ], key[seq_len, head_dim], )5.5 实测数据裸算子 vs 手写 Triton vs torch.compile最后是大家最关心的数字。我用单 batch、长度为 128 的 prefill 和单 token decode 分别做了测试。表格如下方案prefill 延迟 (ms)decode 单 token 延迟 (us)相对 eager 加速比PyTorch eager9.26201.0xtorch.compile6.63801.6x手写 Triton 融合算子5.83201.9xTriton 融合 split-K5.42702.3xTriton 融合 split-K CUDA Graph5.22402.6x所有数字都是同一张 A100 上跑了 200 次取中位数。需要说明的是这只是一个 0.8B 小模型的单机结果不同模型、不同 batch size 下加速比会有差异但趋势是明确的每一步优化都能带来可感知的收益且不影响精度。6. 常见问题与排查技巧实录6.1 问题速查表写这套 kernel 的过程中我踩了不少坑整理成一张速查表方便后面自己看也方便你快速定位。现象可能原因解决办法输出 NaNRoPE 计算中 index 越界或 partner 维度计算错误检查 RoPE 的奇偶维度 partner 逻辑打印中间张量形状输出精度差累加过程中没有转 FP32在 tl.sum、tl.dot 前先把输入转 float32decode 速度没有提升kernel 启动数量过多或 grid 过小用 nsys 分析 kernel 耗时占比考虑用 CUDA Graphsplit-K 结果逐次不一致原子加顺序不确定改用 FP32 原子加或对 split 维做二次归约CUDA Graph 重放时报错捕获时存在未编译 kernel 或动态 shape预热所有 kernel固定输入 shape捕获前 synchronizeattention score 错位QKV 头维度切分错误核对 num_heads、head_dim 与权重 shape 的映射kernel 执行时间反而变长BLOCK_K 过大导致寄存器溢出减小 BLOCK_K 或增加 num_warps显存占用异常KV Cache 预分配过大或激活值未释放监控 torch.cuda.max_memory_allocated6.2 独家经验没有 profile 就没有优化很多新手拿到算子后第一件事就是照着开源实现抄一遍然后发现没快多少。我的建议是动手改之前先跑一遍 nsys 或 ncu看清楚时间到底花在哪个 kernel、哪个阶段、哪块内存拷贝上。我做过一次统计在最初的 PyTorch eager 模式里内存拷贝和 kernel 启动开销占到了总时间的 38%真正执行矩阵乘的时间不到一半。也就是说不管你怎么优化计算只要不减少 kernel 数量和内存拷贝速度永远上不去。这也是我把“21 个算子”“split-K”“CUDA Graph”三件事放在同一篇文章里的原因。它们不是独立优化而是一套组合拳算子融合减少了 kernel 数量和内存往返split-K 提升了并行度CUDA Graph 进一步消灭了启动开销。三者环环相扣缺一个都会让最终效果打折扣。6.3 分享两个容易被忽略的小技巧第一个是“Triton kernel 的边界检查要显式写不能偷懒”。有些 kernel 看着 shape 很规整比如 hidden_size1024BLOCK_K128整除没问题但如果未来切换到别的模型配置可能就出现越界。为了安全所有 tl.load 和 tl.store 我都加了 mask即使当前模型用不到。这个习惯帮我避免了好几次因为切换模型配置导致的 crash。第二个是“把 Triton kernel 的输入输出 shape 尽量固定住”。Triton 对静态 shape 的优化远好于动态 shape尤其是在 autotune 缓存的命中率上。为了让 shape 尽量固定我把序列长度做了 padding 到 128 的倍数decode 阶段 batch 固定为 1。代价是少量显存浪费但换来的是稳定的性能和简单的心智负担我觉得值得。还有一个容易被忽略的问题是“CUDA Graph 捕获期间不要打印张量或调用 item()”。任何 CPU 同步操作都会打断图捕获造成性能骤降甚至报错。如果确实需要在调试时看中间值等 replay 完成后再取。7. 写在最后这套方法还能沿用到哪里这篇文章讲的是 0.8B 模型的 Triton 实现但方法本身是通用的。你手里如果有个 7B 或 13B 模型算子划分思路几乎不变只是矩阵乘的 tile 大小和 split-K 块数需要重新调如果你的模型是 MoE 结构路由相关的算子需要额外处理但注意力部分完全可以复用这套 kernel。我个人实际操作中的体会是写 kernel 最大的障碍不是语法而是心智模型。你得在脑子里同时装着 GPU 的内存层级、SM 调度方式、数据流的依赖关系才能写出既正确又高效的代码。第一次跑通整套 21 个算子的时候我盯着终端里的延迟数字看了很久那种“我能控制 GPU 在干什么”的感觉是调包永远体会不到的。下一篇我会重点拆解注意力算子的 FlashAttention 风格实现以及如何把 online softmax 与 split-K 结合起来在长序列场景下进一步压榨性能。如果你也在手写推理 kernel欢迎在评论区聊聊你遇到的问题一起把坑填平。
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。