Skip to content
KernelIndex
Search⌘K

submission 754244

Navdeep Singh · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-754244?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
56.5µs
#200 of 766
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:1a9e5b0b7690d189f28ca4528129e11b71beea6636b86870a101a083c634f892
license declaredunknown
license concludedunknown
authorsNavdeep Singh
imported2026-08-15

Techniques

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

mmascores = tl.dot(q_lat_fp8, tl.trans(p_lat)) + tl.dot(q_rop_fp8, tl.trans(p_rop))
num-warps = 4num_warps=4,
persistent-kernelnum_splits = tl.num_programs(1)
stages = 3for start_n in tl.range(kv_start, kv_end, BLOCK_N, num_stages=3):

Kernel source

submission.py195 lines
import torch
import triton
import triton.language as tl
from typing import TypeVar

input_t = TypeVar("input_t")
output_t = TypeVar("output_t")

@triton.jit
def mla_stage1_fp8(
    Q, KV_fp8, KV_scale,
    Workspace_V, Workspace_M, Workspace_L, Out,
    kv_indptr, sm_scale,
    stride_qb, stride_qh,
    stride_obs, stride_oh,
    BLOCK_N:   tl.constexpr,
    TILE_SIZE: tl.constexpr,
    WRITE_OUT: tl.constexpr,
):
    LOG2E      = 1.4426950408889634
    batch_id   = tl.program_id(0)
    split_id   = tl.program_id(1)
    num_splits = tl.num_programs(1)

    kv_seq_start = tl.load(kv_indptr + batch_id)
    kv_seq_end   = tl.load(kv_indptr + batch_id + 1)
    
    kv_start     = kv_seq_start + split_id * TILE_SIZE
    kv_end       = tl.minimum(kv_start + TILE_SIZE, kv_seq_end)

    offs_h = tl.arange(0, 16)
    
    if kv_start >= kv_end:
        if not WRITE_OUT:
            ws_ml_off = (batch_id * num_splits + split_id) * 16 + offs_h
            tl.store(Workspace_M + ws_ml_off, -1.0e20)
            tl.store(Workspace_L + ws_ml_off, 0.0)
        return

    global_scale = tl.load(KV_scale).to(tl.float32)

    scale_factor = 200.0  
    outer_scale = (sm_scale * LOG2E * global_scale) / scale_factor
    v_scale = 350.0

    q_base = Q + batch_id * stride_qb
    q_lat_raw = tl.load(q_base + offs_h[:, None] * stride_qh + tl.arange(0, 512)[None, :]).to(tl.float32)
    q_rop_raw = tl.load(q_base + offs_h[:, None] * stride_qh + 512 + tl.arange(0, 64)[None, :]).to(tl.float32)

    fp8_ty = KV_fp8.dtype.element_ty
    q_lat_fp8 = (q_lat_raw * scale_factor).to(fp8_ty)
    q_rop_fp8 = (q_rop_raw * scale_factor).to(fp8_ty)

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

    for start_n in tl.range(kv_start, kv_end, BLOCK_N, num_stages=3):
        offs_n = start_n + tl.arange(0, BLOCK_N)
        mask_n = offs_n < kv_end
        kv_ptr = KV_fp8 + offs_n[:, None] * 576

        p_lat = tl.load(kv_ptr + tl.arange(0, 512)[None, :], mask=mask_n[:, None], other=0.0)
        p_rop = tl.load(kv_ptr + 512 + tl.arange(0, 64)[None, :], mask=mask_n[:, None], other=0.0)

        scores = tl.dot(q_lat_fp8, tl.trans(p_lat)) + tl.dot(q_rop_fp8, tl.trans(p_rop))
        
        scores = scores * outer_scale
        scores = tl.where(mask_n[None, :], scores, -1.0e20)

        m_ij = tl.max(scores, axis=1)
        p    = tl.exp2(scores - m_ij[:, None])
        l_ij = tl.sum(p, axis=1)

        m_next = tl.maximum(m_i, m_ij)
        alpha  = tl.exp2(m_i   - m_next)
        beta   = tl.exp2(m_ij  - m_next)

        # Scale applied *before* accumulation to prevent precision blowout
        p_beta_fp8 = (p * beta[:, None] * v_scale).to(fp8_ty)
        
        acc = acc * alpha[:, None] + tl.dot(p_beta_fp8, p_lat)
        
        l_i = l_i * alpha + l_ij * beta
        m_i = m_next

    acc = (acc * global_scale) / v_scale

    if WRITE_OUT:
        out_ptr = Out + batch_id * stride_obs + offs_h[:, None] * stride_oh + tl.arange(0, 512)[None, :]
        tl.store(out_ptr, (acc / l_i[:, None]).to(tl.bfloat16))
    else:
        ws_ml_off = (batch_id * num_splits + split_id) * 16 + offs_h
        ws_v_off = ((batch_id * num_splits + split_id) * 16 + offs_h[:, None]) * 512
        tl.store(Workspace_V + ws_v_off + tl.arange(0, 512)[None, :], acc.to(tl.bfloat16))
        tl.store(Workspace_M + ws_ml_off, m_i)
        tl.store(Workspace_L + ws_ml_off, l_i)


@triton.jit
def mla_stage2_reduce(
    Workspace_V, Workspace_M, Workspace_L, Out,
    stride_obs, stride_oh,
    num_splits: tl.constexpr, BLOCK_SPLITS: tl.constexpr,
):
    pid      = tl.program_id(0)
    batch_id = pid // 16
    head_id  = pid % 16

    offs_s = tl.arange(0, BLOCK_SPLITS)
    mask_s = offs_s < num_splits

    base_ml = (batch_id * num_splits + offs_s) * 16 + head_id
    m_s = tl.load(Workspace_M + base_ml, mask=mask_s, other=-1.0e20)
    l_s = tl.load(Workspace_L + base_ml, mask=mask_s, other=0.0)

    m_g   = tl.max(m_s, axis=0)
    alpha = tl.exp2(m_s - m_g)
    l_g   = tl.sum(l_s * alpha, axis=0)

    off_v = ((batch_id * num_splits + offs_s) * 16 + head_id) * 512
    v_all = tl.load(Workspace_V + off_v[:, None] + tl.arange(0, 512)[None, :], mask=mask_s[:, None], other=0.0).to(tl.float32)
    v_acc = tl.sum(v_all * alpha[:, None], axis=0)

    tl.store(Out + batch_id * stride_obs + head_id * stride_oh + tl.arange(0, 512), (v_acc / l_g).to(tl.bfloat16))


def _heuristics(bs: int, kv_len: int):
    # The "God Grid" mapping: Target ~256 total blocks globally.
    # We strictly enforce `splits=1` (WRITE_OUT) for large batch sizes.
    
    if bs >= 256: 
        splits = 1    # 256 total blocks
    elif bs >= 64: 
        splits = 4    # 256 total blocks
    elif bs >= 32: 
        splits = 8    # 256 total blocks
    elif bs >= 16: 
        splits = 16   # 256 total blocks
    else: 
        splits = 32   # 128 total blocks (leaves room for scheduler)
        
    block_n = 128
    if kv_len // splits < 128:
        block_n = 64
    if kv_len // splits < 64:
        block_n = 32
        
    tile_size = ((kv_len + splits - 1) // splits)
    # Ensure memory alignment
    tile_size = ((tile_size + block_n - 1) // block_n) * block_n
    
    return int(splits), int(tile_size), int(block_n), 8, 3


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = config["batch_size"]
    kv_len = config.get("kv_seq_len", 8192)

    kv_p_fp8, kv_s_fp8 = kv_data["fp8"] 
    kv_s_fp8 = kv_s_fp8.view(1)

    splits, tile_size, block_n, warps, stages = _heuristics(bs, kv_len)

    out = torch.empty((bs, 16, 512), dtype=torch.bfloat16, device=q.device)
    write_out = (splits == 1)

    ws_v, ws_m, ws_l = None, None, None
    if not write_out:
        ws_v = torch.empty((bs, splits, 16, 512), dtype=torch.bfloat16, device=q.device)
        ws_m = torch.empty((bs, splits, 16), dtype=torch.float32, device=q.device)
        ws_l = torch.empty((bs, splits, 16), dtype=torch.float32, device=q.device)

    mla_stage1_fp8[(bs, splits)](
        q, kv_p_fp8, kv_s_fp8,
        ws_v, ws_m, ws_l, out,
        kv_indptr, config["sm_scale"],
        q.stride(0), q.stride(1),
        out.stride(0), out.stride(1),
        BLOCK_N=block_n, TILE_SIZE=tile_size,
        WRITE_OUT=write_out,
        num_warps=warps, num_stages=stages,
    )

    if not write_out:
        block_splits = triton.next_power_of_2(splits)
        mla_stage2_reduce[(bs * 16,)](
            ws_v, ws_m, ws_l, out,
            out.stride(0), out.stride(1),
            num_splits=splits, BLOCK_SPLITS=block_splits,
            num_warps=4,
        )

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