Skip to content
KernelIndex
Search⌘K

submission 587798

Mihir Shah · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

final.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-587798?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
882.4µs
#725 of 766
2026-03-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:1d12a279dababe01e5a75e94dc34abd20616351da9d2c7e7ebfcb8c07651daf3
license declaredunknown
license concludedunknown
authorsMihir Shah
imported2026-08-26

Techniques

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

fp4kv_p, kv_s = kv_data["mxfp4"]
mmascores = tl.dot(q_le, tl.trans(vke)) + tl.dot(q_lo, tl.trans(vko))
num-warps = 8BLOCK_N=16, TILE_SIZE=1024, num_warps=8, num_stages=1
stages = 1BLOCK_N=16, TILE_SIZE=1024, num_warps=8, num_stages=1
tile-n = 16BLOCK_N=16, TILE_SIZE=1024, num_warps=8, num_stages=1

Kernel source

final.py115 lines
import torch
import triton
import triton.language as tl

NUM_HEADS = 16
V_HEAD_DIM = 512
NUM_KV_SPLITS = 8

@triton.jit
def mla_mxfp4_kernel_seq_split(
    Q, KV_packed, KV_scales, Out, LUT,
    qo_indptr, kv_indptr, sm_scale,
    stride_qb, stride_qh, stride_kvs, stride_kvs_s,
    stride_ob, stride_oh,
    BLOCK_N: tl.constexpr, TILE_SIZE: tl.constexpr,
):
    batch_id = tl.program_id(0)
    tile_id = tl.program_id(1)
    offs_h = tl.arange(0, 16)
    
    kv_start = tl.load(kv_indptr + batch_id) + (tile_id * TILE_SIZE)
    kv_end = tl.minimum(tl.load(kv_indptr + batch_id + 1), kv_start + TILE_SIZE)

    if kv_start >= kv_end:
        return

    # Latent dimensions (512)
    idx_lat_e = tl.arange(0, 256) * 2
    idx_lat_o = idx_lat_e + 1
    # RoPE dimensions (64)
    idx_rop_e = tl.arange(0, 32) * 2 + 512
    idx_rop_o = idx_rop_e + 1
    
    q_base = Q + (batch_id * stride_qb) + (offs_h[:, None] * stride_qh)
    q_le = tl.load(q_base + idx_lat_e[None, :]).to(tl.bfloat16)
    q_lo = tl.load(q_base + idx_lat_o[None, :]).to(tl.bfloat16)
    q_re = tl.load(q_base + idx_rop_e[None, :]).to(tl.bfloat16)
    q_ro = tl.load(q_base + idx_rop_o[None, :]).to(tl.bfloat16)

    acc_e = tl.zeros([16, 256], dtype=tl.float32)
    acc_o = tl.zeros([16, 256], dtype=tl.float32)
    m_i = tl.zeros([16], dtype=tl.float32) - 1.0e20
    l_i = tl.zeros([16], dtype=tl.float32)

    for start_n in range(kv_start, kv_end, BLOCK_N):
        offs_n = start_n + tl.arange(0, BLOCK_N)
        mask_n = offs_n < kv_end

        # Load 576 dims (256 bytes latent + 32 bytes RoPE)
        kv_ptr = KV_packed + (offs_n[:, None] * stride_kvs)
        pkd_l = tl.load(kv_ptr + tl.arange(0, 256)[None, :], mask=mask_n[:, None], other=0)
        pkd_r = tl.load(kv_ptr + 256 + tl.arange(0, 32)[None, :], mask=mask_n[:, None], other=0)
        
        # Load 18 scales (16 latent + 2 RoPE)
        s_ptr = KV_scales + (offs_n[:, None] * stride_kvs_s)
        sl = tl.exp2(tl.load(s_ptr + tl.arange(0, 16)[None, :], mask=mask_n[:, None], other=0).to(tl.float32) - 127.0).to(tl.bfloat16)
        sr = tl.exp2(tl.load(s_ptr + 16 + tl.arange(0, 2)[None, :], mask=mask_n[:, None], other=0).to(tl.float32) - 127.0).to(tl.bfloat16)
        
        sl_b = tl.reshape(tl.broadcast_to(sl[:, :, None], [BLOCK_N, 16, 16]), [BLOCK_N, 256])
        sr_b = tl.reshape(tl.broadcast_to(sr[:, :, None], [BLOCK_N, 2, 16]), [BLOCK_N, 32])

        # Dequantize via LUT
        vke = tl.load(LUT + (pkd_l & 0x0F).to(tl.int32)).to(tl.bfloat16) * sl_b
        vko = tl.load(LUT + ((pkd_l >> 4) & 0x0F).to(tl.int32)).to(tl.bfloat16) * sl_b
        rke = tl.load(LUT + (pkd_r & 0x0F).to(tl.int32)).to(tl.bfloat16) * sr_b
        rko = tl.load(LUT + ((pkd_r >> 4) & 0x0F).to(tl.int32)).to(tl.bfloat16) * sr_b

        # Full 576-dim dot product
        scores = tl.dot(q_le, tl.trans(vke)) + tl.dot(q_lo, tl.trans(vko))
        scores += tl.dot(q_re, tl.trans(rke)) + tl.dot(q_ro, tl.trans(rko))
        
        scores *= sm_scale
        scores = tl.where(mask_n[None, :], scores, -1.0e20)

        m_ij = tl.max(scores, axis=1)
        p = tl.exp(scores - m_ij[:, None])
        l_ij = tl.sum(p, axis=1)
        
        m_next = tl.maximum(m_i, m_ij)
        alpha = tl.exp(m_i - m_next)
        beta = tl.exp(m_ij - m_next)
        
        # Accumulate latent values only (512 dims)
        p_bf = p.to(tl.bfloat16)
        acc_e = acc_e * alpha[:, None] + tl.dot(p_bf, vke) * beta[:, None]
        acc_o = acc_o * alpha[:, None] + tl.dot(p_bf, vko) * beta[:, None]
        
        l_i = l_i * alpha + l_ij * beta
        m_i = m_next

    out_base = Out + (batch_id * 16 * 512) + (offs_h[:, None] * 512)
    tl.store(out_base + idx_lat_e[None, :], (acc_e / l_i[:, None]).to(Out.dtype.element_ty))
    tl.store(out_base + idx_lat_o[None, :],  (acc_o / l_i[:, None]).to(Out.dtype.element_ty))

def custom_kernel(data):
    q, kv_data, qo_indptr, kv_indptr, config = data
    kv_p, kv_s = kv_data["mxfp4"]
    
    # E2M1 Standard LUT
    lut = torch.tensor([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, 
                        -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0], 
                       dtype=torch.float32, device="cuda")
                       
    out = torch.empty((q.shape[0], 16, 512), dtype=torch.bfloat16, device="cuda")
    
    grid = (config["batch_size"], NUM_KV_SPLITS)

    mla_mxfp4_kernel_seq_split[grid](
        q, kv_p.view(torch.uint8), kv_s.view(torch.uint8), out, lut,
        qo_indptr, kv_indptr, config["sm_scale"],
        q.stride(0), q.stride(1), kv_p.stride(0), kv_s.stride(0),
        512 * 16, 512,
        BLOCK_N=16, TILE_SIZE=1024, num_warps=8, num_stages=1
    )
    return out
scrolls · 115 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