Skip to content
KernelIndex
Search⌘K

submission 585650

Arseni Ivanov · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-585650?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
163.9µs
#538 of 766
2026-03-18

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:28dd9afcb5f3ed1e8c7c16fb71adbd51568554aaaf7619d3edce889b41ed95c8
license declaredunknown
license concludedunknown
authorsArseni Ivanov
imported2026-08-26

Techniques

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

mmaqk1 = tl.dot(q1_bf16, tl.trans(v_bf16))
num-warps = 4num_warps=4,
online-softmaxm_new = tl.maximum(m_i, m_ij)
split-kreturn custom_kernel_fp8_splitk(data)
stages = 3num_stages=3

Kernel source

submission_3.py333 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
from task import input_t, output_t

from aiter import dtypes as aiter_dtypes
FP8_DTYPE = aiter_dtypes.fp8

QKV_DTYPE = "fp8"


def custom_kernel(data: input_t) -> output_t:
    """Dispatch to the appropriate kernel based on QKV_DTYPE."""
    if QKV_DTYPE == "fp8":
        qo_indptr = data[2]
        batch_size = qo_indptr.shape[0] - 1
        
        if batch_size > 64:
            return custom_kernel_fp8_nosplit(data)
        else:
            return custom_kernel_fp8_splitk(data)
            
    elif QKV_DTYPE == "bf16":
        return custom_kernel_bf16(data)
    else:
        raise ValueError(f"Invalid QKV_DTYPE: {QKV_DTYPE}")


@triton.jit
def mla_decode_fp8_nosplit_kernel(
    q_ptr, kv_ptr, out_ptr,
    qo_indptr, kv_indptr,
    stride_q_tok, stride_q_h, stride_q_d,
    stride_kv_tok, stride_kv_h, stride_kv_d,
    stride_out_tok, stride_out_h, stride_out_d,
    sm_scale, kv_scale,
    BLOCK_KV: tl.constexpr,
):
    batch_idx = tl.program_id(0)

    q_start = tl.load(qo_indptr + batch_idx)
    kv_start = tl.load(kv_indptr + batch_idx)
    kv_end = tl.load(kv_indptr + batch_idx + 1)
    seq_kv = kv_end - kv_start

    offs_h = tl.arange(0, 16)
    offs_d1 = tl.arange(0, 512)
    offs_d2 = tl.arange(0, 64)

    q_base = q_ptr + q_start * stride_q_tok + offs_h[:, None] * stride_q_h
    q1_ptrs = q_base + offs_d1[None, :] * stride_q_d
    q2_ptrs = q_base + (512 + offs_d2[None, :]) * stride_q_d

    q1_bf16 = (tl.load(q1_ptrs).to(tl.float32) * sm_scale).to(tl.bfloat16)
    q2_bf16 = (tl.load(q2_ptrs).to(tl.float32) * sm_scale).to(tl.bfloat16)

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

    offs_kv = tl.arange(0, BLOCK_KV)
    kv_base = kv_ptr + kv_start * stride_kv_tok

    for start_n in range(0, seq_kv, BLOCK_KV):
        start_n = tl.multiple_of(start_n, BLOCK_KV)
        mask_kv = (start_n + offs_kv) < seq_kv

        curr_kv_ptrs = kv_base + (start_n + offs_kv)[:, None] * stride_kv_tok
        v_ptrs = curr_kv_ptrs + offs_d1[None, :] * stride_kv_d
        k2_ptrs = curr_kv_ptrs + (512 + offs_d2[None, :]) * stride_kv_d

        v_fp8 = tl.load(v_ptrs, mask=mask_kv[:, None], other=0.0)
        k2_fp8 = tl.load(k2_ptrs, mask=mask_kv[:, None], other=0.0)

        v_bf16 = (v_fp8.to(tl.float32) * kv_scale).to(tl.bfloat16)
        k2_bf16 = (k2_fp8.to(tl.float32) * kv_scale).to(tl.bfloat16)

        qk1 = tl.dot(q1_bf16, tl.trans(v_bf16))
        qk2 = tl.dot(q2_bf16, tl.trans(k2_bf16))
        qk = qk1 + qk2

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

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

        l_ij = tl.sum(p, 1)
        l_i = l_i * alpha + l_ij

        acc = acc * alpha[:, None]
        acc += tl.dot(p.to(tl.bfloat16), v_bf16)
        m_i = m_new

    acc = acc / l_i[:, None]
    out_base = out_ptr + q_start * stride_out_tok + offs_h[:, None] * stride_out_h
    out_ptrs = out_base + offs_d1[None, :] * stride_out_d
    tl.store(out_ptrs, acc.to(tl.bfloat16))

def custom_kernel_fp8_nosplit(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    sm_scale = config["sm_scale"]

    kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
    kv_scale_val = kv_scale_fp8.item()

    batch_size = qo_indptr.shape[0] - 1
    total_q = q.shape[0]

    out = torch.empty((total_q, 16, 512), dtype=torch.bfloat16, device=q.device)
    BLOCK_KV = 128

    mla_decode_fp8_nosplit_kernel[(batch_size,)](
        q, kv_buffer_fp8, out,
        qo_indptr, kv_indptr,
        q.stride(0), q.stride(1), q.stride(2),
        kv_buffer_fp8.stride(0), kv_buffer_fp8.stride(1), kv_buffer_fp8.stride(2),
        out.stride(0), out.stride(1), out.stride(2),
        sm_scale, kv_scale_val,
        BLOCK_KV=BLOCK_KV,
    )
    return out


@triton.jit
def mla_decode_fp8_splitk_kernel(
    q_ptr, kv_ptr, 
    workspace_acc, workspace_m, workspace_l,
    out_ptr,
    qo_indptr, kv_indptr,
    stride_q_tok, stride_q_h, stride_q_d,
    stride_kv_tok, stride_kv_h, stride_kv_d,
    stride_out_tok, stride_out_h, stride_out_d,
    combined_scale, kv_scale,
    SPLIT_K: tl.constexpr,
    BLOCK_KV: tl.constexpr,
):
    batch_idx = tl.program_id(0)
    split_idx = tl.program_id(1)

    q_start = tl.load(qo_indptr + batch_idx)
    kv_start = tl.load(kv_indptr + batch_idx)
    kv_end = tl.load(kv_indptr + batch_idx + 1)
    
    seq_kv = kv_end - kv_start
    chunk_size = tl.cdiv(seq_kv, SPLIT_K)
    chunk_start_idx = split_idx * chunk_size
    chunk_end_idx = tl.minimum((split_idx + 1) * chunk_size, seq_kv)

    offs_h = tl.arange(0, 16)
    offs_d1 = tl.arange(0, 512)
    offs_d2 = tl.arange(0, 64)

    q_base = q_ptr + q_start * stride_q_tok + offs_h[:, None] * stride_q_h
    q1_ptrs = q_base + offs_d1[None, :] * stride_q_d
    q2_ptrs = q_base + (512 + offs_d2[None, :]) * stride_q_d

    q1_bf16 = (tl.load(q1_ptrs).to(tl.float32) * combined_scale).to(tl.bfloat16)
    q2_bf16 = (tl.load(q2_ptrs).to(tl.float32) * combined_scale).to(tl.bfloat16)

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

    offs_kv = tl.arange(0, BLOCK_KV)
    kv_base = kv_ptr + kv_start * stride_kv_tok

    for start_n in range(chunk_start_idx, chunk_end_idx, BLOCK_KV):
        start_n = tl.multiple_of(start_n, BLOCK_KV)
        mask_kv = (start_n + offs_kv) < chunk_end_idx

        curr_kv_ptrs = kv_base + (start_n + offs_kv)[:, None] * stride_kv_tok
        v_ptrs = curr_kv_ptrs + offs_d1[None, :] * stride_kv_d
        k2_ptrs = curr_kv_ptrs + (512 + offs_d2[None, :]) * stride_kv_d

        v_fp8 = tl.load(v_ptrs, mask=mask_kv[:, None], other=0.0)
        k2_fp8 = tl.load(k2_ptrs, mask=mask_kv[:, None], other=0.0)

        v_bf16 = v_fp8.to(tl.bfloat16)
        k2_bf16 = k2_fp8.to(tl.bfloat16)

        qk1 = tl.dot(q1_bf16, tl.trans(v_bf16))
        qk2 = tl.dot(q2_bf16, tl.trans(k2_bf16))
        qk = qk1 + qk2
        qk = tl.where(mask_kv[None, :], qk, float("-inf"))

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

        l_ij = tl.sum(p, 1)
        l_i = l_i * alpha + l_ij

        acc = acc * alpha[:, None]
        acc += tl.dot(p.to(tl.bfloat16), v_bf16)
        m_i = m_new

    if SPLIT_K == 1:
        final_out = (acc / l_i[:, None]) * kv_scale
        out_base = out_ptr + q_start * stride_out_tok + offs_h[:, None] * stride_out_h
        out_ptrs = out_base + offs_d1[None, :] * stride_out_d
        tl.store(out_ptrs, final_out.to(tl.bfloat16))
    else:
        ws_idx = batch_idx * SPLIT_K + split_idx
        tl.store(workspace_m + ws_idx * 16 + offs_h, m_i)
        tl.store(workspace_l + ws_idx * 16 + offs_h, l_i)
        ws_acc_base = workspace_acc + ws_idx * 16 * 512 + offs_h[:, None] * 512
        tl.store(ws_acc_base + offs_d1[None, :], acc)


@triton.jit
def mla_decode_reduce_kernel(
    workspace_acc, workspace_m, workspace_l, out_ptr,
    qo_indptr, stride_out_tok, stride_out_h, stride_out_d,
    kv_scale,
    SPLIT_K: tl.constexpr
):
    batch_idx = tl.program_id(0)
    q_start = tl.load(qo_indptr + batch_idx)

    offs_h = tl.arange(0, 16)
    offs_d1 = tl.arange(0, 512)

    m_global = tl.full([16], float("-inf"), dtype=tl.float32)
    l_global = tl.full([16], 0.0, dtype=tl.float32)
    acc_global = tl.zeros([16, 512], dtype=tl.float32)

    for split_idx in range(SPLIT_K):
        ws_idx = batch_idx * SPLIT_K + split_idx
        
        m_j = tl.load(workspace_m + ws_idx * 16 + offs_h)
        l_j = tl.load(workspace_l + ws_idx * 16 + offs_h)
        ws_acc_base = workspace_acc + ws_idx * 16 * 512 + offs_h[:, None] * 512
        acc_j = tl.load(ws_acc_base + offs_d1[None, :])

        m_new = tl.maximum(m_global, m_j)
        alpha_global = tl.exp(m_global - m_new)
        alpha_j = tl.exp(m_j - m_new)
        
        l_global = l_global * alpha_global + l_j * alpha_j
        acc_global = acc_global * alpha_global[:, None] + acc_j * alpha_j[:, None]
        m_global = m_new

    out = (acc_global / l_global[:, None]) * kv_scale
    out_base = out_ptr + q_start * stride_out_tok + offs_h[:, None] * stride_out_h
    tl.store(out_base + offs_d1[None, :], out.to(tl.bfloat16))


def custom_kernel_fp8_splitk(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    sm_scale = config["sm_scale"]

    kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
    kv_scale_val = kv_scale_fp8.item()

    batch_size = qo_indptr.shape[0] - 1
    total_q = q.shape[0]
    combined_scale = sm_scale * kv_scale_val

    target_blocks = 256
    split_k = max(1, min(64, target_blocks // batch_size))

    out = torch.empty((total_q, 16, 512), dtype=torch.bfloat16, device=q.device)

    if split_k > 1:
        workspace_acc = torch.empty((batch_size, split_k, 16, 512), dtype=torch.float32, device=q.device)
        workspace_m = torch.empty((batch_size, split_k, 16), dtype=torch.float32, device=q.device)
        workspace_l = torch.empty((batch_size, split_k, 16), dtype=torch.float32, device=q.device)
    else:
        workspace_acc = q
        workspace_m = q
        workspace_l = q

    BLOCK_KV = 128

    grid_compute = (batch_size, split_k)
    mla_decode_fp8_splitk_kernel[grid_compute](
        q, kv_buffer_fp8, 
        workspace_acc, workspace_m, workspace_l,
        out,
        qo_indptr, kv_indptr,
        q.stride(0), q.stride(1), q.stride(2),
        kv_buffer_fp8.stride(0), kv_buffer_fp8.stride(1), kv_buffer_fp8.stride(2),
        out.stride(0), out.stride(1), out.stride(2),
        combined_scale, kv_scale_val,
        SPLIT_K=split_k,
        BLOCK_KV=BLOCK_KV,
        num_warps=4,
        num_stages=3
    )

    if split_k > 1:
        grid_reduce = (batch_size,)
        mla_decode_reduce_kernel[grid_reduce](
            workspace_acc, workspace_m, workspace_l, out,
            qo_indptr, out.stride(0), out.stride(1), out.stride(2),
            kv_scale_val, SPLIT_K=split_k, num_warps=4
        )

    return out

def custom_kernel_bf16(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    kv_lora_rank = config["kv_lora_rank"]
    sm_scale = config["sm_scale"]
    kv_buffer_bf16 = kv_data["bf16"]
    batch_size = qo_indptr.shape[0] - 1
    out_list = []

    for i in range(batch_size):
        q_s, q_e = int(qo_indptr[i].item()), int(qo_indptr[i + 1].item())
        kv_s, kv_e = int(kv_indptr[i].item()), int(kv_indptr[i + 1].item())

        qi = q[q_s:q_e]                        
        kvc = kv_buffer_bf16[kv_s:kv_e, 0]    
        ki = kvc                               
        vi = kvc[:, :kv_lora_rank]             

        qi_t = qi.float().permute(1, 0, 2) 
        scores = torch.matmul(qi_t * sm_scale, ki.float().T)  
        scores = F.softmax(scores, dim=-1)

        oi = torch.matmul(scores, vi.float())
        oi = oi.permute(1, 0, 2)               
        out_list.append(oi.to(torch.bfloat16))

    return torch.cat(out_list, dim=0)
scrolls · 333 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