Skip to content
KernelIndex
Search⌘K

submission 608517

Amanpreet Singh · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-608517?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
75.9µs
#374 of 766
2026-03-22

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:77930926974fc42e6322126278d340c30fc026e87319c227077e60cdfcbfa5de
license declaredunknown
license concludedunknown
authorsAmanpreet Singh
imported2026-08-26

Techniques

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

mmaqk = tl.dot(q_v_bf16, tl.trans(v_bf16), out_dtype=tl.float32) + tl.dot(q_r_bf16, tl.trans(r_bf16), out_dtype=tl.float32)
num-warps = 4NUM_SPLITS=NS, V_DIM=v_dim, num_warps=4, num_stages=1
persistent-kerneln_splits = tl.num_programs(1)
split-kdef _mla_fp8_mqa_splitk(
stages = 3V_DIM=v_dim, ROPE_DIM=r_dim, BLOCK_KV=B_KV, num_warps=nw, num_stages=3

Kernel source

submission.py171 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t

@triton.jit
def _mla_fp8_mqa_splitk(
    Q, KV_Data, P_Acc, P_Lse,
    qo_indptr, kv_indptr,
    stride_q_b, stride_q_h, stride_q_d,
    stride_kv_b, stride_kv_d,
    stride_pa_b, stride_pa_h, stride_pa_s, stride_pa_d,
    stride_pl_b, stride_pl_h, stride_pl_s,
    KV_Scale, sm_scale,
    V_DIM: tl.constexpr, ROPE_DIM: tl.constexpr, BLOCK_KV: tl.constexpr
):
    b_idx = tl.program_id(0)
    s_idx = tl.program_id(1)
    n_splits = tl.num_programs(1)

    q_st = tl.load(qo_indptr + b_idx)
    kv_st = tl.load(kv_indptr + b_idx)
    kv_end = tl.load(kv_indptr + b_idx + 1)
    seq_len = kv_end - kv_st

    chunk = (seq_len + n_splits - 1) // n_splits
    start_n = s_idx * chunk
    end_n = tl.minimum(start_n + chunk, seq_len)

    offs_h = tl.arange(0, 16)
    offs_v = tl.arange(0, V_DIM)
    offs_r = tl.arange(0, ROPE_DIM) + V_DIM
    offs_kv = tl.arange(0, BLOCK_KV)

    # Natively load scalars straight from VRAM, bypassing Python JIT stalls
    kv_scale_val = tl.load(KV_Scale)
    combined_qk_scale = sm_scale * kv_scale_val

    q_v_ptr = Q + (q_st * stride_q_b) + offs_h[:, None] * stride_q_h + offs_v[None, :] * stride_q_d
    q_r_ptr = Q + (q_st * stride_q_b) + offs_h[:, None] * stride_q_h + offs_r[None, :] * stride_q_d

    # No PyTorch conversion overhead; we load Q natively as BFLOAT16 in the registers
    q_v_bf16 = tl.load(q_v_ptr)
    q_r_bf16 = tl.load(q_r_ptr)

    kv_v_ptr = KV_Data + ((kv_st + start_n) * stride_kv_b) + offs_kv[:, None] * stride_kv_b + offs_v[None, :] * stride_kv_d
    kv_r_ptr = KV_Data + ((kv_st + start_n) * stride_kv_b) + offs_kv[:, None] * stride_kv_b + offs_r[None, :] * stride_kv_d

    m_i = tl.zeros([16], dtype=tl.float32) - float('inf')
    l_i = tl.zeros([16], dtype=tl.float32)
    acc = tl.zeros([16, V_DIM], dtype=tl.float32)

    for n in range(start_n, end_n, BLOCK_KV):
        n_align = tl.multiple_of(n, BLOCK_KV)
        mask = (n_align + offs_kv) < end_n

        v_fp8 = tl.load(kv_v_ptr, mask=mask[:, None], other=0.0)
        r_fp8 = tl.load(kv_r_ptr, mask=mask[:, None], other=0.0)

        v_bf16 = v_fp8.to(tl.bfloat16)
        r_bf16 = r_fp8.to(tl.bfloat16)

        qk = tl.dot(q_v_bf16, tl.trans(v_bf16), out_dtype=tl.float32) + tl.dot(q_r_bf16, tl.trans(r_bf16), out_dtype=tl.float32)
        qk = qk * combined_qk_scale
        qk = tl.where(mask[None, :], qk, -float('inf'))

        m_ij = tl.maximum(m_i, tl.max(qk, axis=1))
        alpha = tl.exp(m_i - m_ij)
        p = tl.exp(qk - m_ij[:, None])

        acc = acc * alpha[:, None] + tl.dot(p.to(tl.bfloat16), v_bf16, out_dtype=tl.float32)
        l_i = l_i * alpha + tl.sum(p, axis=1)
        m_i = m_ij

        kv_v_ptr += BLOCK_KV * stride_kv_b
        kv_r_ptr += BLOCK_KV * stride_kv_b

    acc = tl.where((l_i > 0.0)[:, None], (acc / l_i[:, None]) * kv_scale_val, 0.0)
    lse = m_i + tl.math.log(l_i)
    lse = tl.where(l_i > 0.0, lse, -float('inf'))

    pa_ptr = P_Acc + (q_st * stride_pa_b) + offs_h[:, None] * stride_pa_h + (s_idx * stride_pa_s) + offs_v[None, :] * stride_pa_d
    pl_ptr = P_Lse + (q_st * stride_pl_b) + offs_h * stride_pl_h + (s_idx * stride_pl_s)
    
    tl.store(pa_ptr, acc)
    tl.store(pl_ptr, lse)

@triton.jit
def _mla_mqa_reduce(
    P_Acc, P_Lse, Out, qo_indptr,
    s_pa_b, s_pa_h, s_pa_s, s_pa_d,
    s_pl_b, s_pl_h, s_pl_s,
    s_o_b, s_o_h, s_o_d,
    NUM_SPLITS: tl.constexpr, V_DIM: tl.constexpr
):
    b_idx = tl.program_id(0)
    h_idx = tl.program_id(1)
    q_st = tl.load(qo_indptr + b_idx)
    
    offs_s = tl.arange(0, NUM_SPLITS)
    offs_v = tl.arange(0, V_DIM)

    pl_ptr = P_Lse + (q_st * s_pl_b) + (h_idx * s_pl_h) + offs_s * s_pl_s
    lse = tl.load(pl_ptr)
    
    max_lse = tl.max(lse, axis=0)
    w = tl.exp(lse - max_lse)
    sum_w = tl.sum(w, axis=0)

    pa_ptr = P_Acc + (q_st * s_pa_b) + (h_idx * s_pa_h) + offs_s[:, None] * s_pa_s + offs_v[None, :] * s_pa_d
    acc = tl.load(pa_ptr)
    
    out = tl.sum(acc * w[:, None], axis=0) / sum_w
    
    o_ptr = Out + (q_st * s_o_b) + (h_idx * s_o_h) + offs_v * s_o_d
    tl.store(o_ptr, out.to(Out.dtype.element_ty))

def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    kv_fp8, scale_fp8 = kv_data["fp8"]
    bs = config["batch_size"]
    nh = config["num_heads"]
    v_dim = config["v_head_dim"]
    r_dim = config["qk_rope_head_dim"]
    seq_len = config["kv_seq_len"]
    tq = q.shape[0]

    out = torch.empty((tq, nh, v_dim), dtype=torch.bfloat16, device=q.device)

    # ------------------------------------------------------------------
    # THE OCCUPANCY ENGINE
    # ------------------------------------------------------------------
    if seq_len <= 1024:
        # Dynamically scale splits to wake up all 304 Compute Units
        target_ns = max(1, 256 // bs)
        NS = min(target_ns, 16) # Cap at 16 to keep loop sizes healthy
        B_KV = 64
        nw = 4
    else:
        # UNLEASHED SPLIT-K: Small inner loops = massive speed on big sequences
        NS = 32
        B_KV = 128
        nw = 8

    p_acc = torch.empty((tq, nh, NS, v_dim), dtype=torch.float32, device=q.device)
    p_lse = torch.empty((tq, nh, NS), dtype=torch.float32, device=q.device)

    _mla_fp8_mqa_splitk[(bs, NS)](
        q, kv_fp8, p_acc, p_lse,
        qo_indptr, kv_indptr,
        q.stride(0), q.stride(1), q.stride(2),
        kv_fp8.stride(0), kv_fp8.stride(2),
        p_acc.stride(0), p_acc.stride(1), p_acc.stride(2), p_acc.stride(3),
        p_lse.stride(0), p_lse.stride(1), p_lse.stride(2),
        scale_fp8, config["sm_scale"],
        V_DIM=v_dim, ROPE_DIM=r_dim, BLOCK_KV=B_KV, num_warps=nw, num_stages=3
    )

    if NS == 1:
        out.copy_(p_acc.squeeze(2).to(torch.bfloat16))
    else:
        _mla_mqa_reduce[(bs, nh)](
            p_acc, p_lse, out, qo_indptr,
            p_acc.stride(0), p_acc.stride(1), p_acc.stride(2), p_acc.stride(3),
            p_lse.stride(0), p_lse.stride(1), p_lse.stride(2),
            out.stride(0), out.stride(1), out.stride(2),
            NUM_SPLITS=NS, V_DIM=v_dim, num_warps=4, num_stages=1
        )

      
    return out
scrolls · 171 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