Skip to content
KernelIndex
Search⌘K

submission 634642

Akash Adsare · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submissionv7.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-634642?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
107.6µs
#475 of 766
2026-03-25

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a73f63b1c048ed1c7f3fc33ef4fbbff468373dcef76ce7c0e6fb26b0025b0485
license declaredunknown
license concludedunknown
authorsAkash Adsare
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_bf16, tl.trans(k_e1_bf16))
num-warps = 4num_warps = 4
stages = 2num_stages = 2
tile-n = 64BLOCK_N = 64

Kernel source

submissionv7.py243 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t

@triton.jit
def fp4_to_float_bitwise(v_u8):
    v = v_u8.to(tl.int32)
    v_abs = v & 7
    val_bits = 1056964608 + (v_abs << 22)   

    val_bits = tl.where(v_abs == 1, 1056964608, val_bits)  

    val_bits = tl.where(v_abs == 0, 0, val_bits)            

    val_bits |= (v & 8) << 28                               

    return val_bits.to(tl.float32, bitcast=True)

@triton.jit
def mla_decode_mxfp4_kernel(
    Q, KV_packed, KV_scales,
    kv_indptr,
    Partial_M, Partial_L, Partial_V,
    sm_scale,
    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)
    split_idx = tl.program_id(1)

    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 kv_len <= 0 or 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

    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, :]

    q_e1_bf16 = (tl.load(q_base + h_off * QK_HEAD_DIM + d_256 * 2).to(tl.float32) * sm_scale).to(tl.bfloat16)
    q_o1_bf16 = (tl.load(q_base + h_off * QK_HEAD_DIM + d_256 * 2 + 1).to(tl.float32) * sm_scale).to(tl.bfloat16)
    q_e2_bf16 = (tl.load(q_base + h_off * QK_HEAD_DIM + 512 + d_32 * 2).to(tl.float32) * sm_scale).to(tl.bfloat16)
    q_o2_bf16 = (tl.load(q_base + h_off * QK_HEAD_DIM + 512 + d_32 * 2 + 1).to(tl.float32) * sm_scale).to(tl.bfloat16)

    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)
    acc_v_odd = tl.zeros([NUM_HEADS, 256], dtype=tl.float32)

    PACKED_STRIDE = QK_HEAD_DIM // 2  

    for n_start in range(curr_start, curr_end, BLOCK_N):
        offs = n_start + tl.arange(0, BLOCK_N)
        n_mask = offs < curr_end
        n_off = offs[:, None]

        k_packed_1 = tl.load(
            KV_packed + n_off * PACKED_STRIDE + d_256,
            mask=n_mask[:, None], eviction_policy="evict_first",
        )
        scales_raw_1 = tl.load(
            KV_scales + n_off * NUM_BLOCKS + tl.arange(0, 16)[None, :],
            mask=n_mask[:, None], eviction_policy="evict_first",
        )

        scales_1 = ((scales_raw_1.to(tl.int32) & 0xFF) << 23).to(tl.float32, bitcast=True)

        v_e1 = fp4_to_float_bitwise(k_packed_1 & 0xF)
        k_e1_bf16 = tl.reshape(
            tl.reshape(v_e1, [BLOCK_N, 16, 16]) * scales_1[:, :, None],
            [BLOCK_N, 256]
        ).to(tl.bfloat16)

        v_o1 = fp4_to_float_bitwise((k_packed_1 >> 4) & 0xF)
        k_o1_bf16 = tl.reshape(
            tl.reshape(v_o1, [BLOCK_N, 16, 16]) * scales_1[:, :, None],
            [BLOCK_N, 256]
        ).to(tl.bfloat16)

        k_packed_2 = tl.load(
            KV_packed + n_off * PACKED_STRIDE + 256 + d_32,
            mask=n_mask[:, None], eviction_policy="evict_first",
        )
        scales_raw_2 = tl.load(
            KV_scales + n_off * NUM_BLOCKS + 16 + tl.arange(0, 2)[None, :],
            mask=n_mask[:, None], eviction_policy="evict_first",
        )
        scales_2 = ((scales_raw_2.to(tl.int32) & 0xFF) << 23).to(tl.float32, bitcast=True)

        v_e2 = fp4_to_float_bitwise(k_packed_2 & 0xF)
        k_e2_bf16 = tl.reshape(
            tl.reshape(v_e2, [BLOCK_N, 2, 16]) * scales_2[:, :, None],
            [BLOCK_N, 32]
        ).to(tl.bfloat16)

        v_o2 = fp4_to_float_bitwise((k_packed_2 >> 4) & 0xF)
        k_o2_bf16 = tl.reshape(
            tl.reshape(v_o2, [BLOCK_N, 2, 16]) * scales_2[:, :, None],
            [BLOCK_N, 32]
        ).to(tl.bfloat16)

        s = tl.dot(q_e1_bf16, tl.trans(k_e1_bf16))
        s += tl.dot(q_o1_bf16, tl.trans(k_o1_bf16))
        s += tl.dot(q_e2_bf16, tl.trans(k_e2_bf16))
        s += tl.dot(q_o2_bf16, tl.trans(k_o2_bf16))

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

        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)

        p_bf16 = p.to(tl.bfloat16)
        acc_v_even = acc_v_even * alpha[:, None] + tl.dot(p_bf16, k_e1_bf16)
        acc_v_odd = acc_v_odd * alpha[:, None] + tl.dot(p_bf16, k_o1_bf16)
        m_i = m_next

    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)
    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)

@triton.jit
def mla_reduce_kernel(
    Partial_M, Partial_L, Partial_V,
    Out,
    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
    acc_v_e = tl.zeros([256], dtype=tl.float32)
    acc_v_o = tl.zeros([256], dtype=tl.float32)
    d_256 = tl.arange(0, 256)

    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
        acc_v_e = acc_v_e * alpha_f + tl.load(Partial_V + off_p * V_HEAD_DIM + d_256) * alpha_s
        acc_v_o = acc_v_o * alpha_f + tl.load(Partial_V + off_p * V_HEAD_DIM + 256 + d_256) * alpha_s
        m_final = m_next

    out_ptr = Out + (q_row_idx * NUM_HEADS + h_idx) * V_HEAD_DIM
    inv_l = 1.0 / l_final
    tl.store(out_ptr + d_256 * 2, (acc_v_e * inv_l).to(tl.bfloat16))
    tl.store(out_ptr + d_256 * 2 + 1, (acc_v_o * inv_l).to(tl.bfloat16))

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

    kv_p_u8 = kv_packed.view(torch.uint8)
    kv_s_u8 = kv_scales.view(torch.uint8)

    total_q = q.shape[0]
    NUM_HEADS = 16
    QK_HEAD_DIM = 576
    V_HEAD_DIM = 512
    NUM_BLOCKS = QK_HEAD_DIM // 32  

    batch_size = config["batch_size"]
    kv_seqlen = config["kv_seq_len"]
    BLOCK_N = 64

    if kv_seqlen <= 1024:
        if batch_size <= 4:
            NUM_SPLITS = 8
        elif batch_size <= 32:
            NUM_SPLITS = 16
        elif batch_size <= 64:
            NUM_SPLITS = 4
        else:
            NUM_SPLITS = 2
    else:
        if batch_size <= 4:
            NUM_SPLITS = 32
        elif batch_size <= 32:
            NUM_SPLITS = 8
        elif batch_size <= 64:
            NUM_SPLITS = 8
        else:
            NUM_SPLITS = 8

    max_splits = max(1, kv_seqlen // BLOCK_N)
    NUM_SPLITS = min(NUM_SPLITS, max_splits)

    num_warps = 4
    num_stages = 2

    partial_m = torch.empty(total_q * NUM_HEADS * NUM_SPLITS,
                            dtype=torch.float32, device="cuda")
    partial_l = torch.empty_like(partial_m)
    partial_v = torch.empty(total_q * NUM_HEADS * NUM_SPLITS * V_HEAD_DIM,
                            dtype=torch.float32, device="cuda")
    out = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM),
                      dtype=torch.bfloat16, device="cuda")

    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=num_warps, num_stages=num_stages,
    )

    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=2,
    )
    return out
scrolls · 243 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