Pico-vLLM开发日志 #15 序列内并行的Attention算子优化:Flash Decode方案

今天是久违的Pico-vLLM开发环节(真是阔别已久啊),原因是最近搞到了能够在一段时间内可靠使用的h200x8集群。在上面进行了一些Profiling之后,惊讶的发现Pico-vLLM比起vLLM居然在h200单卡上性能差很多,甚至能差50%以上。在经过性能分析之后,发现主要是之前的短Prompt和短Decoding step总量/在5070卡上的测试中,原本不是瓶颈的Attention Kernel变成了新的瓶颈,甚至把整个Decode step/Prefill step的性能拖累了5~10倍。

这个问题在之前的5070上没有暴露出来,因此就一直没有改进。但在现在的h200机器上,在类似的配置下就暴露的很明显。这实际上是多种因素共同作用的结果,而且这个差距和卡是消费型/计算型有很大的关系。下面详细分析。

Profiling结果和现象解释

在之前的博客当中,我们已经进行过Pico-vLLM的整体性能优化,测得了完整的Profiling结果。为了防止大家忘记,当时的表格是这样的:

Kernel                              时间占比    实例数    单步耗时
──────────────────────────────────────────────────────────────
GEMM (gemvx, 主)                    63.1%      1700     ~6.8ms
GEMM (gemvx, 次)                    22.2%       560     ~2.4ms
cutlass WMMA (gate_up fused)         5.4%       560     ~0.6ms
RMSNorm elementwise                  2.0%      3380     ~0.2ms
BinaryFunctor (residual add)         1.0%      1700     ~0.1ms
CUDAFunctor_add                      0.9%      2240     ~0.1ms
Decode_Paged_GQAAttention            0.9%       560     ~0.1ms
RMSNorm reduce                       0.9%      1140     ~0.1ms
CatArrayBatchedCopy                  0.8%      1120     ~0.08ms
store_kvcache_kernel                 0.2%       560     ~0.02ms
──────────────────────────────────────────────────────────────
GPU 总时间(20步)                  ~214ms     单步 ~10.7ms

当时的结论是,Attention根本啥都不是(只有0.9%),因此Kernel Fusion带来的效果才是最明显的,应该优先做这一类,而Attention做了用处也不大。实际上,在不同length的5070测试情况下,这个论断几乎都是成立的,实测占比几乎没有超过过3%。因此,我当时把它当成了一个固定结论,来作为Profiling的最终结果。当然,事实证明它不是固定结论,至少在更换了GPU之后不一定成立。

然后我们来看在h200单卡上的Profiling结果:

| Time (%) | Total Time (ns) | Instances | Avg (ns) | Med (ns) | Min (ns) | Max (ns) | StdDev (ns) |                                                 Name                                                 |
+----------+-----------------+-----------+----------+----------+----------+----------+-------------+------------------------------------------------------------------------------------------------------+
|     79.6 |       120599594 |       560 | 215356.4 | 215408.0 |   210656 |   220288 |      1671.5 | Decode_Paged_GQAAttention_Kernel                                                                     |
|      6.4 |         9757298 |       560 |  17423.7 |  17344.0 |    16703 |    20096 |       340.3 | nvjet_tst_192x8_64x8_4x1_v_bz_TNT                                                                    |
|      3.7 |         5626849 |       560 |  10047.9 |  10048.0 |     9280 |    11456 |       365.8 | nvjet_tst_64x8_64x16_4x1_v_bz_splitK_TNT                                                             |
|      2.4 |         3636096 |       560 |   6493.0 |   6496.0 |     6080 |     6880 |       124.9 | nvjet_tst_16x64_64x16_4x1_v_bz_TNN                                                                   |
|      1.9 |         2863200 |       560 |   5112.9 |   5088.0 |     4832 |     6112 |       174.6 | nvjet_tst_64x8_64x16_4x1_v_badd_TNT                                                                  |
|      1.5 |         2215232 |        20 | 110761.6 | 110608.0 |   110336 |   111648 |       361.5 | nvjet_tst_384x8_64x4_2x1_v_bz_TNT                                                                    |
|      1.3 |         1959873 |      1140 |   1719.2 |   1696.0 |     1472 |     2112 |       108.9 | _rmsnorm_kernel                                                                                      |
|      1.1 |         1601344 |      1120 |   1429.8 |   1408.0 |     1248 |     2336 |       108.6 | void at::native::vectorized_elementwise_kernel<(int)8, at::native::CUDAFunctor_add<c10::BFloat16>, … |
|      0.8 |         1279808 |       560 |   2285.4 |   2272.0 |     2112 |     3104 |       141.6 | _fused_decode_rope_and_cache_kernel                                                                  |
|      0.6 |          908192 |       560 |   1621.8 |   1600.0 |     1568 |     1792 |        32.1 | void cublasLt::splitKreduce_kernel<(int)32, (int)16, int, float, __nv_bfloat16, float, __nv_bfloat1… |
|      0.5 |          771380 |       560 |   1377.5 |   1375.0 |     1248 |     2240 |        90.5 | _fused_swiglu_kernel                                                                                 |
|      0.1 |          191552 |        20 |   9577.6 |   9584.0 |     9216 |    10272 |       271.0 | void at::native::reduce_kernel<(int)512, (int)1, at::native::ReduceOp<c10::BFloat16, at::native::Ar… |
|      0.1 |           96736 |        40 |   2418.4 |   2432.0 |     2016 |     2752 |       222.4 | void at::native::index_elementwise_kernel<(int)128, (int)4, void at::native::gpu_index_kernel<void … |
|      0.0 |           52160 |        20 |   2608.0 |   2544.0 |     2400 |     2944 |       167.9 | void at::native::<unnamed>::indexSelectSmallIndex<c10::BFloat16, long, unsigned int, (int)2, (int)2… |
|      0.0 |           23904 |        20 |   1195.2 |   1232.0 |     1056 |     1344 |       103.4 | void at::native::vectorized_elementwise_kernel<(int)4, at::native::FillFunctor<int>, std::array<cha… |
+----------+-----------------+-----------+----------+----------+----------+----------+-------------+------------------------------------------------------------------------------------------------------+

聚合改写成人能看懂的表格形式:

kernel占比每实例每步合计(×28 层)
Decode_Paged_GQAAttention_Kernel79.6%215 µs6.0 ms
各 cuBLAS GEMM(合计)~16.5%~1.25 ms
rmsnorm / 残差 / rope / swiglu / 其他~4%~0.16 ms
总 GPU100%~7.4 ms

完全不一样了。GEMM(那一堆nvjet)基本上退居了绝对的次要地位,所有零零散散的dispatch的Kernel加起来也才占据了20%不到的时间占比。然而,Decode_Paged_GQAAttention_Kernel这一个Kernel却单独占比了足足80%。这和5070的情况可以说是截然相反的。也因此,原本和vllm持平的性能,在h200上只有vllm的50%,甚至更低,在极端情况下可以达到只有40%或者35%。

分析和解释

这么大的相对性能差距显然是存在某种系统性因素的,而且很可能不止一个。实际上,笔者本人认为至少从三个角度来解释这个问题。和所有的性能分析模型一样,最主要的因素仍然是显存带宽计算能力,此外还有一个原本在消费级硬件上不是很明显,但在这个配置下变得不可忽略的角度,并发度。在模型完全相同、计算的带宽需求和算力需求完全相同的情况下,实际上是硬件的这两者的不同规格造成了这个差异。

首先列出5070 Laptop和h200的纸面性能的规格差异:

项目RTX 5070 Laptop (8GB)H200 SXMH200 / 5070L
架构Blackwell (5th-gen Tensor Core)Hopper (4th-gen Tensor Core)
GPU dieGB206GH100
SM 数36132~3.67×
CUDA cores4,60816,896~3.67×
Boost clock~2.35 GHz~1.98 GHz
BF16/FP16 Tensor (dense, FP32 acc)~46 TFLOPS~989 TFLOPS~21×
BF16/FP16 Tensor (dense, FP16 acc)~93 TFLOPS~989 TFLOPS~10.6×
BF16/FP16 Tensor (2:4 稀疏)~93 / ~186 TFLOPS~1,979 TFLOPS~10.6× / ~21×
显存8 GB GDDR7141 GB HBM3e~17.6×
显存位宽128-bit6,144-bit (6 × 1024)
显存带宽~512 GB/s4.8 TB/s~9.4×

然后我们来慢慢分析。

显存带宽和计算能力(Ridge Point位置)

这是最明显的区别所在。从表格里可以看到,对于纯BF16/FP16来说,H200的算力是5070 Laptop的21倍,而显存带宽却只有9.4倍。从Roofline模型里可以知道,这就意味着整个h200硬件的Ridge Point(理论上的最优硬件性能利用点)向右偏移了一倍:同样算术强度的程序,在5070上可能是compute bound,在h200上可能就是Memory bound了。而这又进一步意味着在5070上可以被计算给overlap进而掩盖的访存延迟,在h200上会暴露无遗。这里单点就会拉开大约2x的性能差距。不过,它仍然无法完全解释之前的10x性能差距到底是哪儿来的。

并发度

2x偏移的ridge point可以部分的解释Pico-vLLM和vllm的端到端性能的影响,但其实并不能完全解释Attention算子对整个端到端性能差异的影响:毕竟Attention算子本身的性能差异达到了大约10倍甚至更多,如果仅仅是最优算术强度有偏差,无论如何也应该是大约2x的差距,而不是10x。

这部分额外的影响则是并发度和硬件SM数量的不匹配带来的。它的根本原因在于,原本的计算实现没有提供足够多的CTA,而vllm的Flash Decode实现却可以。这一点在5070 Laptop上体现的并不明显:正如前面的规格参数所示,5070 Laptop只有36个SM,而h200却有足足132个。在h200上,原本并发度不足的劣势就被拉的淋漓尽致了。

以我们使用的案例为例,我们有12个Q head,2个KV head。因此,我们launch的Kernel的grid形状就是:[batch_size, Q head]=[1, 12]。也就是说,无论如何,我们总是只能使用固定的12个SM进行运算。而对于vllm来说,一般对于不同情况可以直接用满所有的SM。

这个数字一算出来,结果就很明了了。对于5070 Laptop来说,我们的Pico-vLLM只能利用其33%(1/3=12/36)的性能,而vllm大约是100%。对于h200来说情况就差距更大了:Pico-vLLM只能利用其9.1%(1/11=12/132)的性能,而vllm仍然是100%。这就带来了3倍的相对差异扩大,因此解释了性能差别中的大部分。也正因此,这个性能问题其实最主要和最根本的问题反而是并发度问题,而不是算术强度等等的问题。要解决这个问题就需要发掘其他地方的潜在并发度。实际上,Flash Decode最初的提出也是因为相同的原因而产生的。

之前的实现

之前的代码如下:

  1
  2
  3
  4
  5
  6
  7
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
@triton.jit
def Decode_Paged_GQAAttention_Kernel(
        q                         ,  # (B, n_heads, 1, head_dim)         query,decode 每步只有 1 个 token
        k_cache                   ,  # (num_blocks, n_kv_heads, block_size, head_dim)  全局 K cache
        v_cache                   ,  # (num_blocks, n_kv_heads, block_size, head_dim)  全局 V cache
        block_table               ,  # (B, MAX_BLOCKS_PER_SEQ)  int32,每个请求的物理块 id,-1 表示未分配
        context_lens              ,  # (B,)             int32,每个请求当前的有效 token 数
        scale              ,         # 1.0 / sqrt(head_dim)
        out,
        
        # Meta-parameters
        # 元参数
        MAX_BLOCKS_PER_SEQ: tl.constexpr,  # 启动时固定,不是运行时变量
        BLOCK_SIZE: tl.constexpr,  #
        HEAD_DIM: tl.constexpr,  #
        N_KV_HEAD: tl.constexpr,
        N_HEAD: tl.constexpr,
    ):               # (B, n_heads, 1, head_dim)
    # grid = (B, n_heads)
    # 每个 program 处理一个 (batch, head) 对
    # program_id[0] = batch_idx
    # program_id[1] = head_idx
    
    pid_batch = tl.program_id(axis=0)
    pid_head = tl.program_id(axis=1)
    kv_head_idx = (pid_head // (N_HEAD // N_KV_HEAD))

    m = float('-inf')
    l = float(0)
    o = tl.zeros((HEAD_DIM, ), dtype=tl.float32)
    
    
    q_ptrs = q + pid_batch * (HEAD_DIM * N_HEAD) + pid_head * (HEAD_DIM) + tl.arange(0, HEAD_DIM)
    q_vec = tl.reshape(tl.load(q_ptrs), (HEAD_DIM, 1))

    context_len = tl.load(context_lens + pid_batch)
    if context_len == 0:
        return
    max_block_index = tl.cdiv(context_len, BLOCK_SIZE)  # 向上取整
    offs_kv = tl.arange(0, BLOCK_SIZE * HEAD_DIM)
    
    for block_idx in range(0, max_block_index):
        
        physical_idx = tl.load(block_table + pid_batch * MAX_BLOCKS_PER_SEQ + block_idx)
        physical_idx = tl.maximum(physical_idx, 0).to(tl.int64)
        base = (physical_idx * N_KV_HEAD * BLOCK_SIZE * HEAD_DIM + kv_head_idx* BLOCK_SIZE * HEAD_DIM)

        # 加载时 mask 掉超出 context_len 的 token
        token_start = block_idx * BLOCK_SIZE
        valid_in_block = tl.minimum(BLOCK_SIZE, context_len - token_start)
        kv_token_mask = tl.arange(0, BLOCK_SIZE * HEAD_DIM) < valid_in_block * HEAD_DIM

        k_ptrs = k_cache + base + offs_kv
        # k_block: (block_size, head_dim)
        k_block = tl.load(k_ptrs, mask = kv_token_mask, other=0.0)
        k_block = tl.reshape(k_block, (BLOCK_SIZE, HEAD_DIM))
        v_ptrs = v_cache + base + offs_kv
        v_block = tl.load(v_ptrs, mask = kv_token_mask, other=0.0)
        v_block = tl.reshape(v_block, (BLOCK_SIZE, HEAD_DIM))

        # mask 最后一个 block 的无效 token
        valid = block_idx * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) < context_len
        q_row = tl.reshape(q_vec, (1, HEAD_DIM))           # (1, HEAD_DIM)
        s = tl.sum(k_block * q_row, axis=1)                 # (BLOCK_SIZE,)
        s = s.to(tl.float32) * scale
        s = tl.where(valid, s, float('-inf'))

        m_new = tl.maximum(m, tl.max(s))
        alpha = tl.exp(m - m_new)           # 旧的缩放因子
        p = tl.exp(s - m_new)          # 当前 block 的权重

        l = l * alpha + tl.sum(p)
        p_col = tl.reshape(p, (BLOCK_SIZE, 1))             # (BLOCK_SIZE, 1)
        o = o * alpha + tl.sum(p_col * v_block, axis=0)    # (HEAD_DIM,)
        m = m_new

    o = o / l  # (HEAD_DIM,)
    o_casted = tl.cast(o, out.dtype.element_ty)
    
    tl.store(out + pid_batch * (HEAD_DIM * N_HEAD) + pid_head * (HEAD_DIM) + tl.arange(0, HEAD_DIM), o_casted)

@torch.compiler.disable
def paged_decode_attention(q, k_cache, v_cache, block_table, context_lens,
                           MAX_BLOCKS_PER_SEQ, BLOCK_SIZE=16):
    B, N_HEAD, _, HEAD_DIM = q.shape
    N_KV_HEAD = k_cache.shape[1]
    scale = 1.0 / (HEAD_DIM ** 0.5)

    out = torch.empty(B, N_HEAD, 1, HEAD_DIM, device=q.device, dtype=q.dtype)

    grid = (B, N_HEAD)
    Decode_Paged_GQAAttention_Kernel[grid](
        q, k_cache, v_cache, block_table, context_lens,
        scale, out,
        MAX_BLOCKS_PER_SEQ=MAX_BLOCKS_PER_SEQ,
        BLOCK_SIZE=BLOCK_SIZE,
        HEAD_DIM=HEAD_DIM,
        N_KV_HEAD=N_KV_HEAD,
        N_HEAD=N_HEAD,
    )
    return out

正如同刚刚的分析所说,它的最大问题有两个:一个是并发度极其有限,另一个是在GQA模式下的重复访存。下文会详细说。

Flash Decode的设计模式

Flash Decode的核心就在于挖掘并发度,实际上是一个面向低并发、小batch或者单batch的推理加速方法。对于非常大的Batch来说这通常不会有明显效果(因为都是利用满了整个硬件资源,但Flash Decode还多了个额外的reduction操作),但是对于大多数Batch达不到大规模生产级别的场景来说就有明显的意义了。

具体来说,它通过增加一个在KV序列维度上的均匀切分num_splits(KV 切分数)来挖掘额外的并行度。每个切分的段内先如同其他段不存在一样执行完整的Flash Attention算法流程,然后所有得到的结果再通过一次额外的规约Reduction来得到最终的全局结果。这个思想本身很简单,但要完整的实现一遍还是挺麻烦的,需要写两个Kernel,或者用更加麻烦的原子操作来实现最终的全局规约。

此外,还有一个在GQA中必须进行的设计改动。在之前的经典Flash Attention实现当中,我们会launch一个每个block对每个“逻辑head”进行计算的Kernel。它在逻辑上是正确的,然而也是低效的。它低效的原因在于会把在GQA中被复用的K/V head重复读取,重复次数和KV group size完全相等,或者说等于Q head num/KV head num。这构成了另一个性能下降的原因。

对于5070来说,由于Blackwell架构,其先进、低延迟的L2缓存往往可以部分甚至相当程度上抵消这个问题,从而让这个问题部分地被掩盖。这也相当程度上贡献了Pico-vLLM和vllm在5070上的性能的追平点。但对于h200而言,这个问题会因为较小的L2和较大的L2延迟而体现的更加显著。因此,需要新增一个额外的“head 打包”操作,也即把每个Q head的GEMV变成每个K/V head的,最后一个维度size=group_size的GEMM。如此做,即可让整个操作等效于q_len=group_size而不是q_len=1,从而节省访存、提高算术强度。

不过同时也需要说明的是,它实际上是有代价的,因为它实际上会把grid的配置从[batch_size, Q head, num_splits]=[1, 12, num_splits]变成[batch_size, KV head, num_splits]=[1, 2, num_splits],因此减少了并发度。不过,既然num_splits同时可以提高并发度,这个问题一般是可以被完全弥补的,因为在最后分段结果上reduction的代价比重复访存的代价还是要小多了。

它的源代码大约如下:

  1
  2
  3
  4
  5
  6
  7
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
@triton.jit
def _flash_decode_partial_kernel(
    q,                  # (B, N_HEAD, HEAD_DIM)  (decode: the seq=1 dim collapsed)
    k_cache,            # (num_blocks, N_KV_HEAD, BLOCK_SIZE, HEAD_DIM)
    v_cache,
    block_table,        # (B, MAX_BLOCKS_PER_SEQ) int32
    context_lens,       # (B,) int32
    scale,
    partial_acc,        # (B, N_HEAD, NUM_SPLITS, HEAD_DIM) fp32
    partial_m,          # (B, N_HEAD, NUM_SPLITS) fp32
    partial_l,          # (B, N_HEAD, NUM_SPLITS) fp32
    MAX_BLOCKS_PER_SEQ: tl.constexpr,
    BLOCK_SIZE: tl.constexpr,
    HEAD_DIM: tl.constexpr,
    N_KV_HEAD: tl.constexpr,
    N_HEAD: tl.constexpr,
    GROUP_SIZE: tl.constexpr,
    PADDED_GROUP: tl.constexpr,
    NUM_SPLITS: tl.constexpr,
):
    pid_b = tl.program_id(0)
    pid_kv = tl.program_id(1)
    pid_s = tl.program_id(2)

    ctx = tl.load(context_lens + pid_b)
    q_head0 = pid_kv * GROUP_SIZE

    row = tl.arange(0, PADDED_GROUP)          # (PG,)
    row_mask = row < GROUP_SIZE
    d = tl.arange(0, HEAD_DIM)                 # (HD,)
    qheads = q_head0 + row                     # (PG,) absolute q-head ids

    # partial output pointers for this (b, group, split)
    pm_ptr = partial_m + pid_b * (N_HEAD * NUM_SPLITS) + qheads * NUM_SPLITS + pid_s
    pl_ptr = partial_l + pid_b * (N_HEAD * NUM_SPLITS) + qheads * NUM_SPLITS + pid_s
    pacc_ptr = (partial_acc
                + pid_b * (N_HEAD * NUM_SPLITS * HEAD_DIM)
                + qheads[:, None] * (NUM_SPLITS * HEAD_DIM)
                + pid_s * HEAD_DIM
                + d[None, :])

    # online-softmax state (per padded q-head row)
    m_i = tl.full((PADDED_GROUP,), float("-inf"), dtype=tl.float32)
    l_i = tl.zeros((PADDED_GROUP,), dtype=tl.float32)
    acc = tl.zeros((PADDED_GROUP, HEAD_DIM), dtype=tl.float32)

    # load Q group: (PG, HD)
    q_ptrs = q + pid_b * (N_HEAD * HEAD_DIM) + qheads[:, None] * HEAD_DIM + d[None, :]
    q_tile = tl.load(q_ptrs, mask=row_mask[:, None], other=0.0)

    # this split's block range; empty splits skip the loop and write the
    # sentinel (m=-inf, l=0, acc=0) so the combine step ignores them.
    total_blocks = tl.cdiv(ctx, BLOCK_SIZE)
    blocks_per_split = tl.cdiv(total_blocks, NUM_SPLITS)
    start_block = pid_s * blocks_per_split
    end_block = tl.minimum(start_block + blocks_per_split, total_blocks)

    offs = tl.arange(0, BLOCK_SIZE)
    for blk in range(start_block, end_block):
        phys = tl.load(block_table + pid_b * MAX_BLOCKS_PER_SEQ + blk)
        phys = tl.maximum(phys, 0).to(tl.int64)
        base = phys * (N_KV_HEAD * BLOCK_SIZE * HEAD_DIM) + pid_kv * (BLOCK_SIZE * HEAD_DIM)

        tok_pos = blk * BLOCK_SIZE + offs                 # (BS,)
        kv_valid = tok_pos < ctx
        kv_ptrs = base + offs[:, None] * HEAD_DIM + d[None, :]   # (BS, HD)
        k_blk = tl.load(k_cache + kv_ptrs, mask=kv_valid[:, None], other=0.0)
        v_blk = tl.load(v_cache + kv_ptrs, mask=kv_valid[:, None], other=0.0)

        # S = Q @ K^T : (PG, BS)
        s = tl.dot(q_tile, tl.trans(k_blk)).to(tl.float32) * scale
        s = tl.where(kv_valid[None, :], s, float("-inf"))

        m_new = tl.maximum(m_i, tl.max(s, axis=1))
        alpha = tl.exp(m_i - m_new)
        p = tl.exp(s - m_new[:, None])                    # (PG, BS)
        l_i = l_i * alpha + tl.sum(p, axis=1)
        acc = acc * alpha[:, None] + tl.dot(p.to(v_blk.dtype), v_blk).to(tl.float32)
        m_i = m_new

    tl.store(pm_ptr, m_i, mask=row_mask)
    tl.store(pl_ptr, l_i, mask=row_mask)
    tl.store(pacc_ptr, acc, mask=row_mask[:, None])


@triton.jit
def _flash_decode_combine_kernel(
    partial_acc,        # (B, N_HEAD, NUM_SPLITS, HEAD_DIM) fp32
    partial_m,          # (B, N_HEAD, NUM_SPLITS) fp32
    partial_l,          # (B, N_HEAD, NUM_SPLITS) fp32
    context_lens,       # (B,) int32
    out,                # (B, N_HEAD, HEAD_DIM)
    HEAD_DIM: tl.constexpr,
    N_HEAD: tl.constexpr,
    NUM_SPLITS: tl.constexpr,
):
    pid_b = tl.program_id(0)
    pid_h = tl.program_id(1)
    d = tl.arange(0, HEAD_DIM)
    out_ptr = out + pid_b * (N_HEAD * HEAD_DIM) + pid_h * HEAD_DIM + d

    ctx = tl.load(context_lens + pid_b)
    if ctx == 0:
        tl.store(out_ptr, tl.zeros((HEAD_DIM,), dtype=out.dtype.element_ty))
        return

    s_idx = tl.arange(0, NUM_SPLITS)
    m_s = tl.load(partial_m + pid_b * (N_HEAD * NUM_SPLITS) + pid_h * NUM_SPLITS + s_idx)
    l_s = tl.load(partial_l + pid_b * (N_HEAD * NUM_SPLITS) + pid_h * NUM_SPLITS + s_idx)
    acc_ptrs = (partial_acc
                + pid_b * (N_HEAD * NUM_SPLITS * HEAD_DIM)
                + pid_h * (NUM_SPLITS * HEAD_DIM)
                + s_idx[:, None] * HEAD_DIM
                + d[None, :])
    acc_s = tl.load(acc_ptrs)                            # (NS, HD)

    m_g = tl.max(m_s)                                    # scalar
    factor = tl.exp(m_s - m_g)                           # (NS,)
    l_g = tl.sum(l_s * factor)
    acc_g = tl.sum(acc_s * factor[:, None], axis=0)      # (HD,)
    out_v = acc_g / l_g
    tl.store(out_ptr, out_v.to(out.dtype.element_ty))

在warpper中按顺序依次调用这两个Kernel,并且相应传参即可。同时不得不感叹,Triton真是大幅度的方便了算子的开发。如果是直写cuda的话,可就不是这寥寥一百行可以解决的工程了。说到这个,要不要在有空的时候好好的调研和学习一下cuTile和CUTLASS呢?