Skip to content
KernelIndex
Search⌘K

submission 695464

sunfj · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 289 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-695464?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, int32

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD Instinct MI355X
97.8µs
#457 of 766
2026-04-02

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ede5329a22a86d033bf0c88d6835cc2629556ce1008751b024d663278b0e168a
license declaredunknown
license concludedunknown
authorssunfj
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

autotune@triton.autotune(
mmaqk = tl.dot(q_512, tl.trans(k_512))
num-warps = 8triton.Config({'BLOCK_H': 64, 'BLOCK_N': 64}, num_stages=3, num_warps=8),
split-kkey=['nq', 'num_batches', 'SPLIT_K']
stages = 3triton.Config({'BLOCK_H': 64, 'BLOCK_N': 64}, num_stages=3, num_warps=8),

Kernel source

submission.py289 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

import torch
import triton
import triton.language as tl

class FlashDecodeCache:
    mid_o = None
    mid_m = None
    mid_l = None
    fp32_one = None

# =====================================================================
# Stage 1: 纯血 MLA 极速内核 (消除 V 矩阵冗余加载)
# =====================================================================
@triton.autotune(
    configs=[
        triton.Config({'BLOCK_H': 64, 'BLOCK_N': 64}, num_stages=3, num_warps=8),
        triton.Config({'BLOCK_H': 32, 'BLOCK_N': 128}, num_stages=3, num_warps=4),
        triton.Config({'BLOCK_H': 32, 'BLOCK_N': 64}, num_stages=4, num_warps=4),
        triton.Config({'BLOCK_H': 16, 'BLOCK_N': 128}, num_stages=4, num_warps=4),
    ],
    key=['nq', 'num_batches', 'SPLIT_K']
)
@triton.jit
def mla_decode_stage1(
    Q, KV, Out, mid_O, mid_M, mid_L,
    qo_indptr, kv_indptr, 
    kv_scale_ptr, sm_scale, 
    stride_qz, stride_qh, stride_qd,
    stride_kz, stride_kh, stride_kd,
    # 【核心优化 1】: 彻底删除了 stride_vz, vh, vd,因为不需要了!
    stride_oz, stride_oh, stride_od,
    nq, num_batches,
    SPLIT_K: tl.constexpr,  
    BLOCK_H: tl.constexpr, 
    BLOCK_N: tl.constexpr,
):
    pid_h = tl.program_id(0)
    batch_idx = tl.program_id(1)
    pid_sk = tl.program_id(2)

    if batch_idx >= num_batches: 
        return

    q_idx = tl.load(qo_indptr + batch_idx)
    kv_start_base = tl.load(kv_indptr + batch_idx)
    kv_end_base = tl.load(kv_indptr + batch_idx + 1)

    total_kv = kv_end_base - kv_start_base
    if total_kv <= 0: 
        return

    chunk_size = tl.cdiv(total_kv, SPLIT_K)
    kv_start = kv_start_base + pid_sk * chunk_size
    kv_end = tl.minimum(kv_end_base, kv_start + chunk_size)

    offs_h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H)
    h_mask = offs_h < nq

    offs_dv = tl.arange(0, 512)

    if kv_start >= kv_end:
        if SPLIT_K > 1:
            m_ptrs = mid_M + batch_idx * (nq * SPLIT_K) + offs_h * SPLIT_K + pid_sk
            tl.store(m_ptrs, -float("inf"), mask=h_mask)
            l_ptrs = mid_L + batch_idx * (nq * SPLIT_K) + offs_h * SPLIT_K + pid_sk
            tl.store(l_ptrs, 0.0, mask=h_mask)
            o_ptrs = mid_O + batch_idx * (nq * SPLIT_K * 512) + offs_h[:, None] * (SPLIT_K * 512) + pid_sk * 512 + offs_dv[None, :]
            tl.store(tl.multiple_of(o_ptrs, [1, 16]), 0.0, mask=h_mask[:, None])
        return

    kv_scale = tl.load(kv_scale_ptr)
    combined_scale = kv_scale * sm_scale

    offs_d_512 = tl.arange(0, 512)
    offs_d_64 = tl.arange(512, 576)

    q_ptrs_base_512 = Q + q_idx * stride_qz + offs_h[:, None] * stride_qh + offs_d_512[None, :] * stride_qd
    q_512 = tl.load(tl.multiple_of(q_ptrs_base_512, [1,16]), mask=h_mask[:, None], other=0.0).to(tl.float32)
    q_512 = (q_512 * combined_scale).to(tl.bfloat16)

    q_ptrs_base_64 = Q + q_idx * stride_qz + offs_h[:, None] * stride_qh + offs_d_64[None, :] * stride_qd
    q_64 = tl.load(tl.multiple_of(q_ptrs_base_64, [1,16]), mask=h_mask[:, None], other=0.0).to(tl.float32)
    q_64 = (q_64 * combined_scale).to(tl.bfloat16)

    m_i = tl.zeros([BLOCK_H], dtype=tl.float32) - float("inf")
    l_i = tl.zeros([BLOCK_H], dtype=tl.float32)
    acc = tl.zeros([BLOCK_H, 512], dtype=tl.float32)

    offs_n = tl.arange(0, BLOCK_N)

    for current_n in range(kv_start, kv_end, BLOCK_N):
        n_idx = current_n + offs_n
        kv_mask = n_idx < kv_end

        k_ptrs_512 = KV + n_idx[:, None] * stride_kz + 0 * stride_kh + offs_d_512[None, :] * stride_kd
        k_512 = tl.load(tl.multiple_of(k_ptrs_512, [1,16]), mask=kv_mask[:, None], other=0.0).to(tl.bfloat16)
        qk = tl.dot(q_512, tl.trans(k_512))

        k_ptrs_64 = KV + n_idx[:, None] * stride_kz + 0 * stride_kh + offs_d_64[None, :] * stride_kd
        k_64 = tl.load(tl.multiple_of(k_ptrs_64, [1,16]), mask=kv_mask[:, None], other=0.0).to(tl.bfloat16)
        qk += tl.dot(q_64, tl.trans(k_64))

        qk = tl.where(kv_mask[None, :], qk, float("-inf"))
        qk = tl.where(h_mask[:, None], qk, float("-inf"))

        m_ij = tl.maximum(m_i, tl.max(qk, axis=1))
        m_ij = tl.where(m_ij == float("-inf"), 0.0, m_ij)

        p = tl.exp(qk - m_ij[:, None])
        l_ij = tl.sum(p, axis=1)

        alpha = tl.exp(m_i - m_ij)
        l_i = l_i * alpha + l_ij
        p_bf16 = p.to(tl.bfloat16)

        acc = acc * alpha[:, None]
        
        # =====================================================================
        # 【核心优化 2】: 绝境逢生!不再去显存读取 v_chunk!
        # 在 MLA 中,V 就是 K 的前 512 维!直接复用 SRAM 里的 k_512 进行矩阵乘加!
        # 直接省去 50% 全局显存带宽,耗时瞬间暴跌!
        # =====================================================================
        acc += tl.dot(p_bf16, k_512) 
        
        m_i = m_ij

    if SPLIT_K == 1:
        acc = (acc / l_i[:, None]) * kv_scale
        out_ptrs = Out + q_idx * stride_oz + offs_h[:, None] * stride_oh + offs_dv[None, :] * stride_od
        tl.store(tl.multiple_of(out_ptrs, [1, 16]), acc.to(Out.dtype.element_ty), mask=h_mask[:, None])
    else:
        m_ptrs = mid_M + batch_idx * (nq * SPLIT_K) + offs_h * SPLIT_K + pid_sk
        tl.store(m_ptrs, m_i, mask=h_mask)

        l_ptrs = mid_L + batch_idx * (nq * SPLIT_K) + offs_h * SPLIT_K + pid_sk
        tl.store(l_ptrs, l_i, mask=h_mask)

        o_ptrs = mid_O + batch_idx * (nq * SPLIT_K * 512) + offs_h[:, None] * (SPLIT_K * 512) + pid_sk * 512 + offs_dv[None, :]
        tl.store(tl.multiple_of(o_ptrs, [1, 16]), acc, mask=h_mask[:, None])

# =====================================================================
# Stage 2: 全局归约 (Reduction)
# =====================================================================
@triton.autotune(
    configs=[
        triton.Config({'BLOCK_H': 64}, num_stages=2, num_warps=4),
        triton.Config({'BLOCK_H': 32}, num_stages=2, num_warps=4),
    ],
    key=['nq', 'num_batches']
)
@triton.jit
def mla_decode_stage2(
    mid_O, mid_M, mid_L, Out,
    qo_indptr, kv_indptr, kv_scale_ptr,
    stride_oz, stride_oh, stride_od,
    nq, num_batches,
    SPLIT_K: tl.constexpr,  
    BLOCK_H: tl.constexpr
):
    pid_h = tl.program_id(0)
    batch_idx = tl.program_id(1)

    if batch_idx >= num_batches: return

    offs_h = pid_h * BLOCK_H + tl.arange(0, BLOCK_H)
    h_mask = offs_h < nq
    offs_dv = tl.arange(0, 512)

    kv_start = tl.load(kv_indptr + batch_idx)
    kv_end = tl.load(kv_indptr + batch_idx + 1)
    if kv_start >= kv_end:
        q_idx = tl.load(qo_indptr + batch_idx)
        out_ptrs = Out + q_idx * stride_oz + offs_h[:, None] * stride_oh + offs_dv[None, :] * stride_od
        tl.store(tl.multiple_of(out_ptrs, [1,16]), 0.0, mask=h_mask[:, None])
        return

    kv_scale = tl.load(kv_scale_ptr)

    global_m = tl.zeros([BLOCK_H], dtype=tl.float32) - float("inf")
    global_l = tl.zeros([BLOCK_H], dtype=tl.float32)
    acc = tl.zeros([BLOCK_H, 512], dtype=tl.float32)

    for sk in range(SPLIT_K):
        m_ptrs = mid_M + batch_idx * (nq * SPLIT_K) + offs_h * SPLIT_K + sk
        m_sk = tl.load(m_ptrs, mask=h_mask, other=-float("inf"))
        global_m = tl.maximum(global_m, m_sk)

    global_m = tl.where(global_m == -float("inf"), 0.0, global_m)

    for sk in range(SPLIT_K):
        m_ptrs = mid_M + batch_idx * (nq * SPLIT_K) + offs_h * SPLIT_K + sk
        m_sk = tl.load(m_ptrs, mask=h_mask, other=-float("inf"))

        l_ptrs = mid_L + batch_idx * (nq * SPLIT_K) + offs_h * SPLIT_K + sk
        l_sk = tl.load(l_ptrs, mask=h_mask, other=0.0)

        alpha = tl.exp(m_sk - global_m)
        global_l += l_sk * alpha

        o_ptrs = mid_O + batch_idx * (nq * SPLIT_K * 512) + offs_h[:, None] * (SPLIT_K * 512) + sk * 512 + offs_dv[None, :]
        o_vals = tl.load(tl.multiple_of(o_ptrs, [1, 16]), mask=h_mask[:, None], other=0.0)

        acc += o_vals * alpha[:, None]

    acc = (acc / global_l[:, None]) * kv_scale

    q_idx = tl.load(qo_indptr + batch_idx)
    out_ptrs = Out + q_idx * stride_oz + offs_h[:, None] * stride_oh + offs_dv[None, :] * stride_od
    tl.store(tl.multiple_of(out_ptrs, [1,16]), acc.to(Out.dtype.element_ty), mask=h_mask[:, None])


def mla_decode_triton(q: torch.Tensor, kv_input: torch.Tensor, qo_indptr: torch.Tensor, kv_indptr: torch.Tensor, config: dict, kv_scale_tensor: torch.Tensor):
    batch_size = config["batch_size"]
    nq = config["num_heads"]
    dq = config["qk_head_dim"]  
    dv = config["v_head_dim"]   
    sm_scale = 1.0 / (dq ** 0.5)

    total_q = q.size(0)
    outputs = torch.empty((total_q, nq, dv), device=q.device, dtype=torch.bfloat16)

    # =====================================================================
    # 并发激进嗅探器:哪怕只有 1024 的序列,也强行切割保证占满 GPU!
    # =====================================================================
    total_kv = kv_input.shape[0]
    avg_seq_len = total_kv // max(1, batch_size)
    
    # 调大目标块数,强行触发 MI355X 并发
    target_blocks = 512 
    base_blocks = max(1, batch_size * max(1, nq // 64))
    
    # 只要序列长于 128,就允许切割!
    max_split = max(1, avg_seq_len // 128) 
    desired_split = triton.cdiv(target_blocks, base_blocks)
    
    # 动态确定最佳分割度
    SPLIT_K = min(16, min(max_split, desired_split))
    MAX_SPLIT_K = 16 

    if SPLIT_K > 1:
        if FlashDecodeCache.mid_o is None or FlashDecodeCache.mid_o.shape[0] < batch_size or FlashDecodeCache.mid_o.shape[2] < MAX_SPLIT_K:
            FlashDecodeCache.mid_o = torch.zeros((batch_size, nq, MAX_SPLIT_K, dv), dtype=torch.float32, device=q.device)
            FlashDecodeCache.mid_m = torch.full((batch_size, nq, MAX_SPLIT_K), float("-inf"), dtype=torch.float32, device=q.device)
            FlashDecodeCache.mid_l = torch.zeros((batch_size, nq, MAX_SPLIT_K), dtype=torch.float32, device=q.device)

    grid1 = lambda META: (triton.cdiv(nq, META['BLOCK_H']), batch_size, SPLIT_K)
    mla_decode_stage1[grid1](
        q, kv_input, outputs, FlashDecodeCache.mid_o, FlashDecodeCache.mid_m, FlashDecodeCache.mid_l,
        qo_indptr, kv_indptr, kv_scale_tensor, sm_scale, 
        q.stride(0), q.stride(1), q.stride(2),
        kv_input.stride(0), kv_input.stride(1), kv_input.stride(2),
        # 移除了 kv_input 第三组 stride,内核签名已同步更新
        outputs.stride(0), outputs.stride(1), outputs.stride(2),
        nq, batch_size,
        SPLIT_K=SPLIT_K,
    )

    if SPLIT_K > 1:
        grid2 = lambda META: (triton.cdiv(nq, META['BLOCK_H']), batch_size)
        mla_decode_stage2[grid2](
            FlashDecodeCache.mid_o, FlashDecodeCache.mid_m, FlashDecodeCache.mid_l, outputs,
            qo_indptr, kv_indptr, kv_scale_tensor,
            outputs.stride(0), outputs.stride(1), outputs.stride(2),
            nq, batch_size,
            SPLIT_K=SPLIT_K,
        )

    return outputs


def custom_kernel(data):
    if FlashDecodeCache.fp32_one is None or FlashDecodeCache.fp32_one.device != data[0].device:
        FlashDecodeCache.fp32_one = torch.tensor([1.0], dtype=torch.float32, device=data[0].device)

    q, kv_data, qo_indptr, kv_indptr, config = data

    if "fp8" in kv_data:
        kv_input, kv_scale_tensor = kv_data["fp8"]
        if kv_scale_tensor is None:
            kv_scale_tensor = FlashDecodeCache.fp32_one
    else:
        kv_input = kv_data.get("bf16", kv_data[list(kv_data.keys())[0]])
        if isinstance(kv_input, tuple): kv_input = kv_input[0]
        kv_scale_tensor = FlashDecodeCache.fp32_one

    return mla_decode_triton(q, kv_input, qo_indptr, kv_indptr, config, kv_scale_tensor)
scrolls · 289 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON