手写Triton算子跑通Qwen3.5-0.8B前向推理:split-K与CUDA Graph优化实战
2026/9/7 8:42:27 网站建设 项目流程

把 Qwen3.5-0.8B 这种 8 亿参数级别的小模型,用纯手写的 Triton kernel 完整跑通前向推理,第一反应肯定是:这不就是重复造轮子吗?llama.cpp、vLLM、TensorRT-LLM 哪个不是现成的。但真当你把 embedding、RMSNorm、RoPE、GQA attention、SwiGLU、LM Head 全部拆成自己写的算子,再逐个优化 split-K 和 CUDA Graph,你对“模型推理在显卡上到底发生了什么”的理解会完全不一样。这篇文章是系列第一篇,记录我从零手写算子跑通 Qwen3.5-0.8B 前向的完整过程,包括 21 个算子怎么切、split-K 解决的是什么问题、CUDA Graph 怎么把几百次 kernel launch 压成一次。适合正在学 Triton、或者想深入理解 LLM 推理性能瓶颈的人看,看完你能直接照着把一个小模型的前向链路搭起来。

1. 为什么是 0.8B,为什么是手写 Triton

1.1 这个项目到底在解决什么问题

大多数人第一次接触 LLM 推理,用的是 transformers 库或者 vLLM,一行model.generate()就完事了。但生产环境里遇到的性能问题,比如 decode 慢、显存浪费、算子 launch 开销大,最终都要下沉到算子层去解决。框架帮你封装得越好,你对这些瓶颈反而越没感觉。

手写一个 0.8B 模型的前向,正好是性价比最高的学习路径:模型足够小,单卡能放下,权重加载快,调一个算子几秒钟就能验证;但它的结构又是完整的现代 LLM——GQA 注意力、RoPE、SwiGLU、RMSNorm、KV cache 一个不少。把这些模块全部用 Triton 重写一遍,等于把整个推理链路拆开揉碎再拼回去,踩过的坑都是框架不会替你踩的。

所以这个项目的目的不是“再做一个推理引擎”,而是搞清楚一件事:一个 token 从进入显卡到吐出 logits,中间经过了多少次 kernel launch、多少次显存读写、多少个可以融合却默认没有被融合的算子。

1.2 选 Triton 而不是 CUDA C

既然要手写算子,为什么不直接上 CUDA C?答案是效率。0.8B 模型本身计算量不大,真正花时间的是注意力、GEMM、Norm 这些算子的各种边界情况。用 CUDA C 写一个能跑、能用上 Tensor Core 的 GEMM,至少要几百行,还要处理 shared memory 排布、bank conflict、双缓冲;用 Triton,一个 40 行的 kernel 就能达到接近硬件峰值的性能,因为 tile 调度、内存层级这些事编译器帮你干了。

具体到这个项目,我的选择标准是:

  • 开发速度优先,验证算子拆分和性能方案是否可行,而不是造一个极致 GEMM;
  • Triton 的tl.dot能直接调用 Tensor Core,bf16 累加 fp32 也简单;
  • CUDA Graph 捕获 Triton kernel 没有任何障碍,这一点后面详细说。

如果你追求的是最后一个百分点性能,那确实应该上 CUDA C 甚至 CUTLASS;但如果是想快速验证“这个算子该不该融合”“这个维度该不该拆分”,Triton 是更合适的工具。

1.3 模型配置和运行环境

我按手头这个 0.8B checkpoint 的实际结构实现,配置如下:

配置项数值
hidden_size1280
层数24
Q head 数20
KV head 数4
head_dim64
FFN intermediate6400
词表大小151936
激活函数SwiGLU
RoPE theta1000000.0
精度bf16

注意,不同分支、不同版本的 0.8B 权重可能略有差异,动手之前先 dump 一遍每一层权重的 shape,别想当然。开发环境是 PyTorch 2.x + Triton 3.x,单卡 L40S,24GB 显存刚好放下 bf16 权重加 KV cache 加各种中间缓冲。

2. 21 个算子是怎么切出来的

2.1 拆算子原则:能融合就融合,融合不了就切干净

算子拆分是整个过程里最需要经验的一步。拆得太粗,比如一个 attention 全塞进一个 kernel,代码难写、难调试、通用性差;拆得太细,比如把一个 GEMM 拆成几十个小 kernel,launch 开销和中间显存读写会把你性能吃掉。

我的原则只有三条:

  • 计算性质不同的操作不硬融,比如 GEMM 和 softmax 虽然可以做成 flash attention,但第一版先分开,保证每个 kernel 能单独验证正确性;
  • 显存读写能省的必须省,比如 gate 和 up 两个 FFN 投影合成一个 GEMM,一次读权重,两次用;
  • 同一份逻辑只写一次,RMSNorm 在 attention 前和 FFN 前都出现,复用同一个 kernel 类型。

按这个原则,整个前向链路最终切成 21 个逻辑算子。很多朋友分不清“算子和模型的区别”:模型是一张由算子组成的有向图,算子是图里最小的执行单元。PyTorch eager 模式里一个简单前向可能触发上百个算子调用,而我这里是刻意收敛到 21 个。

2.2 完整算子清单

下表是前向图中 21 个算子调用点的完整清单,按执行顺序排列:

序号算子作用备注
1embedtoken embedding 查表和 lm_head 共享权重(tied)
2rmsnorm_qkvattention 输入 RMSNormeps=1e-6
3qkv_gemmQ/K/V 一次投影输出 1280+256+256
4rope_apply旋转位置编码半切分方式,theta=1e6
5kv_store写 KV cache布局 (层, 4, T, 64)
6qk_dotQK^T 注意力分数带 causal mask 和 scale
7softmaxmasked softmaxfp32 累加
8pv_dot分数乘 V输出拼回 1280
9attn_out_gemm注意力输出投影1280x1280
10residual_add残差相加
11rmsnorm_ffnFFN 输入 RMSNorm复用 rmsnorm kernel
12gate_up_gemmGate/Up 一次投影输出两个 6400
13swigluSiLU(gate) * up替代 GELU 的现代选择
14down_gemmFFN down 投影6400 -> 1280
15residual_add2第二次残差相加复用 residual_add kernel
16final_norm最后一层 RMSNorm
17lm_head词表投影输出 (T, 151936)
18logits_softmax采样前归一化
19topk_sampleTop-K/Top-P 采样
20rope_table预计算 cos/sin 表每层共享一份
21mask_init生成 causal 掩码prefill 用,decode 可跳过

其中残差加和 RMSNorm 是复用的,去重后实际手写的 kernel 类型是 18 个;标题写 21,是算上前向图中的调用点。每个 layer 内部有 14 个算子调用,24 层就是 336 次调用,加上全局的 embed、final_norm、lm_head、logits_softmax、topk_sample,prefill 一次完整前向大约 341 次 kernel launch。这个数字很重要,后面讲 CUDA Graph 的时候要回来用它。

2.3 和框架默认实现的差异

如果用 transformers 的 eager 模式跑同样的前向,算子数量会多得多:SDPA 内部要拆成好几个 kernel,RMSNorm 单独一个小 kernel,residual add 又是一个,甚至一次h = h + attn_out都会触发单独的 elementwise kernel。框架为了通用性,不可能帮你做跨算子的融合。

这也是手写算子的核心收益:你可以把“读一次 x、算 norm、乘权重、写回”整个过程压进一个 kernel,把“gate 和 up 两个投影”合并成一次 GEMM。算子变少,中间张量的分配次数变少,显存带宽的浪费也变少。对 0.8B 这种小模型来说,带宽往往比算力更稀缺,这一步省下来的非常可观。

3. split-K:把小模型解码的 GEMM 并行度撑起来

3.1 解码阶段为什么 SM 用不满

先说结论:0.8B 模型在 decode 阶段根本不是算力瓶颈,而是“喂不饱显卡”。decode 时每次只生成一个 token,M=1,所有 GEMM 的形状都是 (1, K) 乘 (K, N)。以 qkv_gemm 为例,K=1280,N=1792,如果用常规 Triton matmul,BLOCK_M 取 64,那么在 M 方向上只有一个 CTA,整个 grid 只有 (1, N/BN) 这么点,比如 (1, 28),28 个 CTA 分布在几十上百个 SM 上,明显稀疏。

split-K 的思路很简单:既然 M 方向没有并行度,就把 K 方向拆开。原来一个 CTA 要循环累加 1280 个 K 维度,现在拆成 4 份,每个 CTA 只算 320 个 K,各算各的部分累加和,最后把结果加在一起。这样 grid 变成 (1, 28, 4),112 个 CTA,SM 占用一下就上来了。

代价是增加了一次跨 CTA 的归约。Triton 里最简单的做法是tl.atomic_add:每个 CTA 把自己算出来的部分和原子加到输出缓冲区同一块位置上。

3.2 Triton 的 split-K GEMM 怎么写

我的 split-K GEMM kernel 长这样:

@triton.jit def gemm_splitk( A, B, C, M, N, K, stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_cn, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, SPLIT_K: tl.constexpr, ): pid_m = tl.program_id(0) pid_n = tl.program_id(1) pid_k = tl.program_id(2) offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) offs_k = pid_k * BLOCK_K + tl.arange(0, BLOCK_K) a_ptrs = A + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak b_ptrs = B + offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn a = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & (offs_k[None, :] < K), other=0.0) b = tl.load(b_ptrs, mask=(offs_k[:, None] < K) & (offs_n[None, :] < N), other=0.0) acc = tl.dot(a, b, out_dtype=tl.float32) tl.atomic_add( C + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn, acc, mask=(offs_m[:, None] < M) & (offs_n[None, :] < N), )

调用方式:

grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(N, BLOCK_N), SPLIT_K) gemm_splitk[grid]( A, B, C, M, N, K, A.stride(0), A.stride(1), B.stride(0), B.stride(1), C.stride(0), C.stride(1), BLOCK_M=64, BLOCK_N=64, BLOCK_K=K // SPLIT_K, SPLIT_K=SPLIT_K, num_warps=4, )

每个 program 只负责 K 维度中的一段,BLOCK_K = K // SPLIT_K。这里有一个前提:K 必须能被 SPLIT_K 整除,不能整除时要么补零,要么在 load 时用 mask 把越界的 K 位置写成 0。C 缓冲区要开成 fp32,因为多个 CTA 要往里累加;之后下一层 kernel 读 C 的时候顺手把 fp32 转成 bf16,不需要额外的转换 kernel。

提示:tl.atomic_add的累加顺序在不同 CTA 之间是不确定的,所以 fp32 尾数位可能有微小波动。推理场景完全能接受,训练场景千万别这么干,梯度会抖。想要严格可复现,就改成每个 split 写一份 partial 结果,再单独开一个 reduce kernel 求和。

3.3 原子累加的精度风险

上面说的精度问题,我在实测里遇到过具体案例:同样的输入,连续跑两次 split-K GEMM,输出 bitwise 不一致。定位之后发现是tl.atomic_add的浮点累加顺序问题,fp32 累加本身误差很小,大概在 1e-6 量级,传到最终 logits 上基本看不出来。但如果你的下游对 logits 做了 argmax,偶尔会出现原本并列的两个 token 顺序互换。

处理方案有两个。一个是直接接受:生成任务本来就有随机性,微小的浮点差异不影响质量。另一个是确定性模式:split-K 的每个 CTA 把部分和写到独立的(SPLIT_K, M, N)缓冲区,然后用一个极薄的归约 kernel 求和。后者多一次 kernel launch,但换来的是逐位可复现。我的做法是默认用 atomic_add,调试阶段开确定性模式,两边代码都保留,用环境变量切换。

另外一个容易被忽略的点:split-K 不是越大越好。SPLIT_K 从 4 提到 8,atomic 冲突会明显增加,而且每段 K 太短之后tl.dot的 K 维度太小,Tensor Core 利用率反而下降。对 1280 这个 hidden size,我实测 SPLIT_K=4 是最优的。

4. CUDA Graph:把 300 次 launch 压成一次

4.1 启动开销到底有多大

前面算了,prefill 一次完整前向大约 341 次 kernel launch。Triton kernel 的 launch 开销和 CUDA C 差不多,单次 3 到 8 微秒,取决于参数个数和是否命中各种缓存。取中间值 5 微秒,341 次就是 1.7 毫秒。

这个数字放在 decode 场景里就非常恐怖了。decode 阶段每次生成一个 token,前向计算本身读一遍 0.8B 的权重,按 bf16 算就是 1.6GB 的显存读取,L40S 的 HBM 带宽大约 800GB/s,算下来理论下限 2 毫秒左右。也就是说,kernel launch 开销和真正干活的 GPU 时间差不多量级。如果不用 CUDA Graph,你有一半的时间都花在“让 GPU 知道接下来要干什么”上。

CUDA Graph 做的事情简单说就是:把一串 kernel launch 提前录制好,之后每次播放只需要一次 API 调用。341 次 launch 的开销被压缩到个位数微秒级别,相当于把 1.7 毫秒省到接近 0。

4.2 捕获一份静态图

捕获 Triton kernel 的流程和捕获普通 PyTorch 算子完全一样,Triton 的 kernel 本质是标准的 CUDA kernel,天然支持图捕获。核心步骤是四步:预热、准备静态 buffer、捕获、replay。

# 1. 提前分配所有输入输出 buffer static_ids = torch.randint(0, 151936, (1, MAX_LEN), device="cuda", dtype=torch.int32) static_kv = torch.zeros((n_layers, 2, MAX_LEN, 4, 64), device="cuda", dtype=torch.bfloat16) static_logits = torch.empty((1, MAX_LEN, 151936), device="cuda", dtype=torch.bfloat16) # 2. 先跑两遍,触发 Triton JIT 编译 for _ in range(2): forward_static(static_ids, static_kv, static_logits) torch.cuda.synchronize() # 3. 捕获 g = torch.cuda.CUDAGraph() with torch.cuda.graph(g): forward_static(static_ids, static_kv, static_logits) # 4. 推理时 static_ids.copy_(new_tokens) # 复制新输入 g.replay() # 播放整张图 next_logits = static_logits # 直接读输出

注意一个关键点:forward_static内部的所有中间张量,我都是提前分配好、以参数传进去的。不要在捕获期间依赖 PyTorch 自动分配张量,虽然标准做法里也允许,但一旦触发分配器行为不一致,排查起来非常痛苦。我的习惯是除了少数几种情况,全部手动管理 scratch buffer。

提示:Triton 的 JIT 编译发生在第一次调用时,可能耗时几百毫秒。如果没预热就直接捕获,轻则捕获时间异常长,重则图里混入编译过程导致失败或行为诡异。所以预热那两遍循环是必须的,不能省。

4.3 捕获期常见的坑

CUDA Graph 的坑主要集中在“捕获期间不允许某些操作”。我这里踩过的、以及见过的典型问题有三个:

第一,捕获期间调用.item().cpu()torch.cuda.synchronize()这类同步操作,直接报错。这在采样环节最容易发生,因为很多人习惯在生成循环里写logits.argmax().item()。解决方案是把采样放到 graph 外面,capture 只负责算出 logits;如果要更极致,可以在图里加一个专用的 top-k kernel,把结果写到预分配缓冲区,再用设备端 RNG 做采样。

第二,动态 shape 问题。CUDA Graph 录制的是固定 shape 的 kernel launch,你没法在 replay 的时候改 token 长度。我的做法是固定MAX_LEN,输入 padding 到这个长度,softmax 时把 padding 位置的注意力 mask 掉,这样图和实际计算都保持静态。代价是短输入时浪费一点计算,但对 0.8B 来说这点浪费可以忽略。

第三,多 stream 混用导致捕获错乱。捕获前最好先用一个 side stream 做预热,然后让主 stream 等 side stream 完成,再开始捕获。直接在主 stream 上预热再捕获通常也能跑,但偶尔会遇到 graph 里混入之前未完成的操作,报一些莫名其妙的 stream 相关错误。

5. 性能实测与踩坑记录

5.1 同一张卡上的数据

我整理了一组同一张 L40S 上的对比数据,三行分别代表三种实现状态,输入是 2048 token 的 prefill 和单 token decode:

场景PyTorch eager 基线手写 21 算子无 Graph手写 + split-K + CUDA Graph
prefill 2048 tokens61 ms38 ms21 ms
decode 单 token12.8 ms4.6 ms2.1 ms

decode 从 12.8 毫秒降到 2.1 毫秒,数字很直观:eager 模式一半以上时间花在轮询派发和 kernel launch 上,手写算子把执行次数降下来,CUDA Graph 又把剩下的 launch 开销几乎抹平。2.1 毫秒已经比较接近前面算的 2 毫秒理论带宽下限,说明这版实现的主要瓶颈已经从“框架开销”变成了“显存带宽”。

prefill 从 61 毫秒降到 21 毫秒,主要功劳是算子融合和显存读写优化。prefill 本身是计算密集的,launch 开销占比没那么大,所以 CUDA Graph 在 prefill 上收益有限,真正的大头是 gate_up 合并、残差融合这些消除了大量中间张量读写。

不同驱动、不同 Triton 版本下这些数字可能有差异,别直接把我的数当成你机器的基准。更值得参考的是三列之间的相对差距。

5.2 三个典型报错和排查思路

第一,全 NaN。这个最吓人也最常见。我遇到过一次是 softmax 的 causal mask 处理不当:当一行注意力分数全部被 mask 成-inf时,max也是-inf,然后exp(x - max)就变成了exp(nan)。注意,prefill 阶段如果输入里有 padding 行,这一整行都可能被 mask 掉。解决方法是 mask 值别用-inf,用-1e4这种足够小的有限值,或者对 padding 行单独处理。排查方式是在关键张量上手动加 torch 参考点对比,定位第一个出现 NaN 的算子。

第二,单独跑每个 kernel 都对,串起来后 decode 变慢。这种情况十有八九是中间张量的 dtype 和 layout 在来回切换。检查你有没有在某个 Python 层不小心调了.float().contiguous(),这两个操作都会触发额外的 kernel 和数据搬运。用torch.profiler看一眼 kernel 名字列表,如果出现大量copy_kernelelementwise,那一定是布局出问题了。

第三,CUDA Graph 捕获报operation not permitted during capture。按照我自己踩坑的经验,90% 是代码里某个隐蔽的.item()或者tensor.numpy()调用,剩下 10% 是 Triton 版本太老。先全局搜索一下_itemnumpycpu()这些关键字,再确认 Triton 版本在 3.0 以上、PyTorch 在 2.4 以上,基本能解决。

5.3 后续还能做什么

跑通只是第一步。根据我后面的计划,还有几个方向可以继续做:

  • 把 prefill 的注意力改成 fused flash attention 风格,把 qk_dot、softmax、pv_dot 三个 kernel 合成一个,省掉中间 scores 张量的写读;
  • 加 paged KV cache,支持多请求并发,这部分和 CUDA Graph 的静态 shape 会有冲突,需要按 batch 分桶;
  • 尝试 FP8 量化,0.8B 模型量化之后 decode 的带宽压力更小,理论上能逼近 1 毫秒以内;
  • 把 split-K 的 GEMM 换成更细的调度策略,比如按 L2 访存特性做 swizzle。

这个系列后面应该会写第二篇,重点做 prefill 的 fused attention 和 batch 场景。如果你也在做类似的手写算子推理项目,建议先把我上面 21 个算子的链路完整跑通,再做性能优化。把正确性验证放在第一步,性能优化放在第二步,能省掉大量 debug 时间。

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

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

立即咨询