Skip to content
KernelIndex
Search⌘K

submission 694148

sky · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-694148?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
4.83ms
#757 of 766
2026-04-02

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9cb458629cfc09d68b99cb59c6c68d80a4e9a8a7ab80b790be782b02cd4f5113
license declaredunknown
license concludedunknown
authorssky
imported2026-08-26

Techniques

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

fp4if "mxfp4" in kv_data:

Kernel source

submission3.py112 lines
import torch
import math

# -----------------------------
# Constants
# -----------------------------
NUM_HEADS = 16
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / math.sqrt(QK_HEAD_DIM)


# -----------------------------
# Dequant
# -----------------------------
def dequant_fp8(x, scale):
    return x.to(torch.bfloat16) * scale


def dequant_mxfp4(fp4_data, scale_e8m0):
    from aiter.utility.fp4_utils import mxfp4_to_f32, e8m0_to_f32

    B, M, N2 = fp4_data.shape
    N = N2 * 2
    rows = B * M

    x = mxfp4_to_f32(fp4_data.reshape(rows, N2))
    scales = e8m0_to_f32(scale_e8m0)

    block = 32
    num_blocks = N // block
    scales = scales[:rows, :num_blocks]

    x = x.view(rows, num_blocks, block)
    x = x * scales.unsqueeze(-1)

    return x.view(B, M, N).to(torch.bfloat16)


# -----------------------------
# FAST batched attention
# -----------------------------
def fast_attention(q, k, v):
    """
    q: (B, 16, 576)
    k: (B, L, 576)
    v: (B, L, 512)
    """

    # (B, 16, L)
    scores = torch.matmul(q.to(torch.float32), k.transpose(1, 2).to(torch.float32))
    scores *= SM_SCALE

    probs = torch.softmax(scores, dim=-1)

    # (B, 16, 512)
    out = torch.matmul(probs, v.to(torch.float32))

    return out.to(torch.bfloat16)


# -----------------------------
# Main kernel (vectorized)
# -----------------------------
def mla_decode_kernel(data):
    q, kv_data, qo_indptr, kv_indptr, config = data
    device = q.device

    batch_size = qo_indptr.shape[0] - 1
    q = q.view(batch_size, NUM_HEADS, QK_HEAD_DIM)

    # -------------------------
    # Resolve KV
    # -------------------------
    if "mxfp4" in kv_data:
        kv_buffer, scale = kv_data["mxfp4"]
        kv_full = dequant_mxfp4(kv_buffer, scale)

    elif "fp8" in kv_data:
        kv_buffer, scale = kv_data["fp8"]
        kv_full = dequant_fp8(kv_buffer, scale)

    else:
        kv_full = kv_data["bf16"]

    kv_full = kv_full.squeeze(1)  # (total_kv, 576)

    outputs = []

    # -------------------------
    # Batch (still segmented, but vectorized inside)
    # -------------------------
    for b in range(batch_size):
        kv_start = kv_indptr[b].item()
        kv_end = kv_indptr[b + 1].item()

        k = kv_full[kv_start:kv_end, :QK_HEAD_DIM].unsqueeze(0)   # (1, L, 576)
        v = kv_full[kv_start:kv_end, :V_HEAD_DIM].unsqueeze(0)    # (1, L, 512)

        qb = q[b].unsqueeze(0)  # (1, 16, 576)

        out = fast_attention(qb, k, v)  # (1, 16, 512)
        outputs.append(out)

    return torch.cat(outputs, dim=0)
    

# -----------------------------
# Entry point
# -----------------------------
def custom_kernel(data):
    return mla_decode_kernel(data)
scrolls · 112 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