Skip to content
KernelIndex
Search⌘K

submission 608589

akasha_08267 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:31d456cdc6596a584eb08021b1855404b1b509a0b1456f6bd4d1ed9a787899b7
license declaredunknown
license concludedunknown
authorsakasha_08267
imported2026-08-26

Techniques

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

fp4kv_packed, kv_scales = kv_data["mxfp4"]
mmas = tl.dot(q_e1, tl.trans(k_e1))
num-warps = 4num_warps=4,
stages = 2num_stages=2,
tile-n = 64BLOCK_N = 64

Kernel source

submissionv2.py316 lines
import math
import torch
import triton
import triton.language as tl
from task import input_t, output_t

# ---------------------------------------------------------------------------
# DeepSeek R1 MLA constants
# ---------------------------------------------------------------------------
NUM_HEADS        = 16
NUM_KV_HEADS     = 1
KV_LORA_RANK     = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM      = KV_LORA_RANK + QK_ROPE_HEAD_DIM   # 576
V_HEAD_DIM       = KV_LORA_RANK                      # 512
SM_SCALE         = 1.0 / (QK_HEAD_DIM ** 0.5)

# MXFP4 specifics
BLOCK_D     = 32
NUM_BLOCKS  = QK_HEAD_DIM // BLOCK_D                 # 18
NUM_SPLITS  = 32
BLOCK_N     = 64


# ---------------------------------------------------------------------------
# 4-bit OCP-MX to float32 conversion
# ---------------------------------------------------------------------------
@triton.jit
def fp4_to_float(v):
    """Convert 4-bit E2M1 (OCP MX) values to float32."""
    sign = (v >> 3) & 1
    exp  = (v >> 1) & 3
    mant = v & 1

    # Denormal: exponent == 0
    denorm = mant.to(tl.float32) * 0.5
    # Normal: (1 + mant*0.5) * 2^(exp-1)
    norm   = (1.0 + mant.to(tl.float32) * 0.5) * tl.exp2((exp - 1).to(tl.float32))
    val    = tl.where(exp == 0, denorm, norm)
    return tl.where(sign > 0, -val, val)


# ---------------------------------------------------------------------------
# Main MXFP4 decode kernel (persistent KV splits)
# ---------------------------------------------------------------------------
@triton.jit
def mla_decode_mxfp4_kernel(
    Q,                 # (total_q, NUM_HEADS, QK_HEAD_DIM) bf16
    KV_packed,         # (total_kv, 1, QK_HEAD_DIM//2) uint8 / fp4x2
    KV_scales,         # (total_kv, 1, NUM_BLOCKS) uint8 (E8M0)
    kv_indptr,         # (batch_size+1,) int32
    Partial_M,         # (total_q * NUM_HEADS * NUM_SPLITS) float32
    Partial_L,         # (total_q * NUM_HEADS * NUM_SPLITS) float32
    Partial_V,         # (total_q * NUM_HEADS * NUM_SPLITS * V_HEAD_DIM) float32
    sm_scale,          # float32
    NUM_HEADS: tl.constexpr,
    QK_HEAD_DIM: tl.constexpr,
    V_HEAD_DIM: tl.constexpr,
    NUM_SPLITS: tl.constexpr,
    BLOCK_N: tl.constexpr,
    NUM_BLOCKS: tl.constexpr,
):
    q_row_idx = tl.program_id(0)   # which query row
    split_idx = tl.program_id(1)   # which split within that row

    # -----------------------------
    # 1. Determine KV window
    # -----------------------------
    kv_start = tl.load(kv_indptr + q_row_idx)
    kv_end   = tl.load(kv_indptr + q_row_idx + 1)
    kv_len   = kv_end - kv_start

    kv_per_split = (kv_len + NUM_SPLITS - 1) // NUM_SPLITS
    curr_start   = kv_start + split_idx * kv_per_split
    curr_end     = tl.minimum(curr_start + kv_per_split, kv_end)

    # If no KV tokens for this row or this split, write neutral partials and exit.
    if (kv_len <= 0) | (curr_start >= curr_end):
        h_range = tl.arange(0, NUM_HEADS)
        off_p   = (q_row_idx * NUM_HEADS + h_range) * NUM_SPLITS + split_idx
        tl.store(Partial_M + off_p, tl.full([NUM_HEADS], -float('inf'), dtype=tl.float32))
        tl.store(Partial_L + off_p, tl.zeros([NUM_HEADS], dtype=tl.float32))
        return

    # -----------------------------
    # 2. Load Q for all heads into registers
    # -----------------------------
    q_base = Q + q_row_idx * NUM_HEADS * QK_HEAD_DIM
    h_off  = tl.arange(0, NUM_HEADS)[:, None]
    d_256  = tl.arange(0, 256)[None, :]
    d_32   = tl.arange(0, 32)[None, :]

    # 576 dims = 2*256 + 2*32 (even/odd)
    q_e1 = tl.load(q_base + h_off * QK_HEAD_DIM + d_256 * 2).to(tl.float32)
    q_o1 = tl.load(q_base + h_off * QK_HEAD_DIM + d_256 * 2 + 1).to(tl.float32)
    q_e2 = tl.load(q_base + h_off * QK_HEAD_DIM + (256 + d_32) * 2).to(tl.float32)
    q_o2 = tl.load(q_base + h_off * QK_HEAD_DIM + (256 + d_32) * 2 + 1).to(tl.float32)

    # -----------------------------
    # 3. Online-softmax accumulators
    # -----------------------------
    m_i = tl.full([NUM_HEADS], -float('inf'), dtype=tl.float32)
    l_i = tl.zeros([NUM_HEADS], dtype=tl.float32)

    acc_v_even = tl.zeros([NUM_HEADS, 256], dtype=tl.float32)  # first 256 of 512
    acc_v_odd  = tl.zeros([NUM_HEADS, 256], dtype=tl.float32)  # last 256 of 512

    # -----------------------------
    # 4. Iterate over KV tokens in BLOCK_N tiles
    # -----------------------------
    for n_start in range(curr_start, curr_end, BLOCK_N):
        n_offsets = n_start + tl.arange(0, BLOCK_N)
        n_mask    = n_offsets < curr_end

        # 4.1 Load packed MXFP4 KV
        # Each row has QK_HEAD_DIM//2 bytes (two 4-bit values per byte).
        kv_row_stride_bytes = QK_HEAD_DIM // 2

        # First 256 dims use the first 256 bytes (512 values)
        k_packed_1 = tl.load(
            KV_packed + n_offsets[:, None] * kv_row_stride_bytes + d_256,
            mask=n_mask[:, None],
            other=0,
            eviction_policy="evict_first"
        )
        # Last 64 dims use the next 32 bytes
        k_packed_2 = tl.load(
            KV_packed + n_offsets[:, None] * kv_row_stride_bytes + 256 + d_32,
            mask=n_mask[:, None],
            other=0,
            eviction_policy="evict_first"
        )

        # 4.2 Load block scales (E8M0, 18 blocks total)
        #   First 8 blocks cover the first 256 dims, the next 2 for the 64 dims.
        scales_raw_1 = tl.load(
            KV_scales + n_offsets[:, None] * NUM_BLOCKS + tl.arange(0, 8)[None, :],
            mask=n_mask[:, None],
            other=0,
            eviction_policy="evict_first"
        )
        scales_raw_2 = tl.load(
            KV_scales + n_offsets[:, None] * NUM_BLOCKS + 8 + tl.arange(0, 2)[None, :],
            mask=n_mask[:, None],
            other=0,
            eviction_policy="evict_first"
        )

        # Convert E8M0 to float scalars
        scales_1 = tl.exp2(scales_raw_1.to(tl.float32) - 127.0)  # (BLOCK_N, 8)
        scales_2 = tl.exp2(scales_raw_2.to(tl.float32) - 127.0)  # (BLOCK_N, 2)

        # 4.3 Unpack MXFP4 → float32 and apply scales
        # nibble low/high
        v_e1 = fp4_to_float((k_packed_1 & 0xF).to(tl.int32))
        v_o1 = fp4_to_float(((k_packed_1 >> 4) & 0xF).to(tl.int32))
        v_e2 = fp4_to_float((k_packed_2 & 0xF).to(tl.int32))
        v_o2 = fp4_to_float(((k_packed_2 >> 4) & 0xF).to(tl.int32))

        # First 256 dims: 8 blocks x 32 features = 256.
        k_e1 = tl.reshape(
            tl.reshape(v_e1, [BLOCK_N, 8, 32]) * scales_1[:, :, None],
            [BLOCK_N, 256]
        )
        k_o1 = tl.reshape(
            tl.reshape(v_o1, [BLOCK_N, 8, 32]) * scales_1[:, :, None],
            [BLOCK_N, 256]
        )

        # Last 64 dims: 2 blocks x 16 features = 32 per half.
        k_e2 = tl.reshape(
            tl.reshape(v_e2, [BLOCK_N, 2, 16]) * scales_2[:, :, None],
            [BLOCK_N, 32]
        )
        k_o2 = tl.reshape(
            tl.reshape(v_o2, [BLOCK_N, 2, 16]) * scales_2[:, :, None],
            [BLOCK_N, 32]
        )

        # 4.4 Compute scores for each head: s = q·k^T
        #   q_e1: (16, 256), k_e1: (BLOCK_N, 256) → s1: (16, BLOCK_N)
        s = tl.dot(q_e1, tl.trans(k_e1))
        s += tl.dot(q_e2, tl.trans(k_e2))
        s += tl.dot(q_o1, tl.trans(k_o1))
        s += tl.dot(q_o2, tl.trans(k_o2))

        s = s * sm_scale
        s = tl.where(n_mask[None, :], s, -float('inf'))

        # 4.5 Online softmax update
        m_ij   = tl.max(s, axis=1)
        m_next = tl.maximum(m_i, m_ij)

        alpha = tl.exp(m_i - m_next)
        p     = tl.exp(s - m_next[:, None])
        l_i   = l_i * alpha + tl.sum(p, axis=1)

        # Value accumulation uses only first 512 dims → first 256 even + 256 odd.
        acc_v_even = acc_v_even * alpha[:, None] + tl.dot(p, k_e1)
        acc_v_odd  = acc_v_odd  * alpha[:, None] + tl.dot(p, k_o1)

        m_i = m_next

    # -----------------------------
    # 5. Write partial results
    # -----------------------------
    h_range = tl.arange(0, NUM_HEADS)
    off_p   = (q_row_idx * NUM_HEADS + h_range) * NUM_SPLITS + split_idx

    tl.store(Partial_M + off_p, m_i)
    tl.store(Partial_L + off_p, l_i)

    # Flatten V-head dimension: each head/split has V_HEAD_DIM = 512 values.
    tl.store(Partial_V + off_p[:, None] * V_HEAD_DIM + d_256,        acc_v_even)
    tl.store(Partial_V + off_p[:, None] * V_HEAD_DIM + 256 + d_256,  acc_v_odd)


# ---------------------------------------------------------------------------
# Reduction kernel over splits
# ---------------------------------------------------------------------------
@triton.jit
def mla_reduce_kernel(
    Partial_M, Partial_L, Partial_V,
    Out,            # (total_q, NUM_HEADS, V_HEAD_DIM) bf16
    NUM_HEADS: tl.constexpr,
    V_HEAD_DIM: tl.constexpr,
    NUM_SPLITS: tl.constexpr,
):
    q_row_idx = tl.program_id(0)
    h_idx     = tl.program_id(1)

    m_final = -float('inf')
    l_final = 0.0
    d_256   = tl.arange(0, 256)

    acc_v_e = tl.zeros([256], dtype=tl.float32)
    acc_v_o = tl.zeros([256], dtype=tl.float32)

    # Reduce across splits
    for s in range(NUM_SPLITS):
        off_p = (q_row_idx * NUM_HEADS + h_idx) * NUM_SPLITS + s

        m_s = tl.load(Partial_M + off_p)
        l_s = tl.load(Partial_L + off_p)

        m_next   = tl.maximum(m_final, m_s)
        alpha_f  = tl.exp(m_final - m_next)
        alpha_s  = tl.exp(m_s      - m_next)

        l_final  = l_final * alpha_f + l_s * alpha_s

        v_e_s = tl.load(Partial_V + off_p * V_HEAD_DIM + d_256)
        v_o_s = tl.load(Partial_V + off_p * V_HEAD_DIM + 256 + d_256)

        acc_v_e = acc_v_e * alpha_f + v_e_s * alpha_s
        acc_v_o = acc_v_o * alpha_f + v_o_s * alpha_s

        m_final = m_next

    # Normalize
    acc_v_e = acc_v_e / l_final
    acc_v_o = acc_v_o / l_final

    # Write out
    out_ptr = Out + (q_row_idx * NUM_HEADS + h_idx) * V_HEAD_DIM
    tl.store(out_ptr + d_256 * 2,       acc_v_e.to(tl.bfloat16))
    tl.store(out_ptr + d_256 * 2 + 1,   acc_v_o.to(tl.bfloat16))


# ---------------------------------------------------------------------------
# Python wrapper matching (q, kv_data, qo_indptr, kv_indptr, config) -> out
# ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    kv_packed, kv_scales = kv_data["mxfp4"]

    # Ensure uint8 for Triton loads
    kv_p_u8 = kv_packed.view(torch.uint8)
    kv_s_u8 = kv_scales.view(torch.uint8)

    total_q = q.shape[0]
    device  = q.device

    # Allocate partial buffers
    partial_m = torch.empty(total_q * NUM_HEADS * NUM_SPLITS, dtype=torch.float32, device=device)
    partial_l = torch.empty_like(partial_m)
    partial_v = torch.empty(total_q * NUM_HEADS * NUM_SPLITS * V_HEAD_DIM, dtype=torch.float32, device=device)

    out = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device)

    # Launch decode kernel: grid = (total_q, NUM_SPLITS)
    mla_decode_mxfp4_kernel[(total_q, NUM_SPLITS)](
        q, kv_p_u8, kv_s_u8, kv_indptr,
        partial_m, partial_l, partial_v,
        float(config["sm_scale"]),
        NUM_HEADS=NUM_HEADS,
        QK_HEAD_DIM=QK_HEAD_DIM,
        V_HEAD_DIM=V_HEAD_DIM,
        NUM_SPLITS=NUM_SPLITS,
        BLOCK_N=BLOCK_N,
        NUM_BLOCKS=NUM_BLOCKS,
        num_warps=4,
        num_stages=2,
    )

    # Launch reduction kernel: grid = (total_q, NUM_HEADS)
    mla_reduce_kernel[(total_q, NUM_HEADS)](
        partial_m, partial_l, partial_v, out,
        NUM_HEADS=NUM_HEADS,
        V_HEAD_DIM=V_HEAD_DIM,
        NUM_SPLITS=NUM_SPLITS,
        num_warps=4,
        num_stages=2,
    )

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