Skip to content
KernelIndex
Search⌘K

submission 661945

mega-dmitriy · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v104_pro1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-661945?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
44.5µs
#136 of 766
2026-03-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8c2b8a48e00bd1f5b84b52e3959990b36e724956dc81f04211f8d183d0132d61
license declaredunknown
license concludedunknown
authorsmega-dmitriy
imported2026-08-15

Techniques

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

fp8q_lat = tl.load(Q + q_idx * stride_q_tok + offs_h[:, None] * stride_q_head + offs_dk_lat[None, :]).to(tl.float8e4nv)
mmascores = (tl.dot(q_lat, tl.trans(k_lat_fp8)) + tl.dot(q_rope, tl.trans(k_rope_fp8))).to(tl.float32) * score_scale
online-softmaxm_new = tl.maximum(m_i, m_ij)
persistent-kernelout_base = (split_id * tl.num_programs(1) + batch_id) * NUM_HEADS_CONST
split-ksplit_kv_start = kv_start + split_id * TILES_PER_SPLIT * BLOCK_KV
stages = 2TILES_TOTAL=kv_seq_len // BKV, num_stages=2, allow_flush_denorm=True)

Kernel source

submission_v104_pro1.py298 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
V104_pro1: Rec 5 + Rec 4 — AITER num_kv_splits sweep (24 instead of 32) + exp2.
- Try num_kv_splits=24 for bs=256/kv=8k (researcher suggests 24 as first sweep point)
- exp2 softmax + allow_flush_denorm on all Triton kernels
- Based on v69_c (current best 46.69us)
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t

from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
LOG2E = 1.4426950408889634
SM_SCALE_LOG2E = SM_SCALE * LOG2E
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8


@triton.jit
def _fused_fp8qk_singlepass(
    Q, KV_FP8, KV_SCALE, O,
    kv_indptr, qo_indptr, sm_scale_log2e,
    stride_q_tok, stride_q_head, stride_kv_tok,
    BLOCK_KV: tl.constexpr, BLOCK_DV: tl.constexpr,
    NUM_HEADS_CONST: tl.constexpr, TILES_TOTAL: tl.constexpr,
):
    batch_id = tl.program_id(0)
    kv_start = tl.load(kv_indptr + batch_id)
    q_idx = tl.load(qo_indptr + batch_id)
    offs_h = tl.arange(0, NUM_HEADS_CONST)
    offs_dv = tl.arange(0, BLOCK_DV)
    offs_dk_lat = tl.arange(0, 512)
    offs_dk_rope = tl.arange(0, 64)
    offs_kv = tl.arange(0, BLOCK_KV)
    q_lat = tl.load(Q + q_idx * stride_q_tok + offs_h[:, None] * stride_q_head + offs_dk_lat[None, :]).to(tl.float8e4nv)
    q_rope = tl.load(Q + q_idx * stride_q_tok + offs_h[:, None] * stride_q_head + 512 + offs_dk_rope[None, :]).to(tl.float8e4nv)
    kv_s = tl.load(KV_SCALE)
    score_scale = sm_scale_log2e * kv_s
    m_i = tl.full([NUM_HEADS_CONST], value=float("-inf"), dtype=tl.float32)
    l_i = tl.zeros([NUM_HEADS_CONST], dtype=tl.float32)
    acc = tl.zeros([NUM_HEADS_CONST, BLOCK_DV], dtype=tl.float32)
    for tile_idx in range(TILES_TOTAL):
        kv_offset = kv_start + tile_idx * BLOCK_KV
        k_lat_fp8 = tl.load(KV_FP8 + (kv_offset + offs_kv[:, None]) * stride_kv_tok + offs_dk_lat[None, :])
        k_rope_fp8 = tl.load(KV_FP8 + (kv_offset + offs_kv[:, None]) * stride_kv_tok + 512 + offs_dk_rope[None, :])
        scores = (tl.dot(q_lat, tl.trans(k_lat_fp8)) + tl.dot(q_rope, tl.trans(k_rope_fp8))).to(tl.float32) * score_scale
        m_ij = tl.max(scores, axis=1)
        m_new = tl.maximum(m_i, m_ij)
        alpha = tl.math.exp2(m_i - m_new)
        p = tl.math.exp2(scores - m_new[:, None])
        l_i = alpha * l_i + tl.sum(p, axis=1)
        k_lat_bf16 = k_lat_fp8.to(tl.bfloat16)
        acc = acc * alpha[:, None] + tl.dot(p.to(tl.bfloat16), k_lat_bf16)
        m_i = m_new
    acc = (acc * kv_s) / tl.maximum(l_i[:, None], 1e-12)
    out_ptr = O + q_idx * NUM_HEADS_CONST * BLOCK_DV
    tl.store(out_ptr + offs_h[:, None] * BLOCK_DV + offs_dv[None, :], acc.to(tl.bfloat16))


@triton.jit
def _flash_decode_fp8qk_stage1(
    Q, KV_FP8, KV_SCALE, partial_O, partial_m, partial_l,
    kv_indptr, qo_indptr, sm_scale_log2e,
    stride_q_tok, stride_q_head, stride_kv_tok,
    NUM_SPLITS: tl.constexpr, BLOCK_KV: tl.constexpr,
    BLOCK_DV: tl.constexpr, NUM_HEADS_CONST: tl.constexpr,
    TILES_PER_SPLIT: tl.constexpr,
):
    split_id = tl.program_id(0)
    batch_id = tl.program_id(1)
    kv_start = tl.load(kv_indptr + batch_id)
    q_idx = tl.load(qo_indptr + batch_id)
    split_kv_start = kv_start + split_id * TILES_PER_SPLIT * BLOCK_KV
    out_base = (split_id * tl.num_programs(1) + batch_id) * NUM_HEADS_CONST
    offs_h = tl.arange(0, NUM_HEADS_CONST)
    offs_dv = tl.arange(0, BLOCK_DV)
    offs_dk_lat = tl.arange(0, 512)
    offs_dk_rope = tl.arange(0, 64)
    offs_kv = tl.arange(0, BLOCK_KV)
    q_lat = tl.load(Q + q_idx * stride_q_tok + offs_h[:, None] * stride_q_head + offs_dk_lat[None, :]).to(tl.float8e4nv)
    q_rope = tl.load(Q + q_idx * stride_q_tok + offs_h[:, None] * stride_q_head + 512 + offs_dk_rope[None, :]).to(tl.float8e4nv)
    kv_s = tl.load(KV_SCALE)
    score_scale = sm_scale_log2e * kv_s
    m_i = tl.full([NUM_HEADS_CONST], value=float("-inf"), dtype=tl.float32)
    l_i = tl.zeros([NUM_HEADS_CONST], dtype=tl.float32)
    acc = tl.zeros([NUM_HEADS_CONST, BLOCK_DV], dtype=tl.float32)
    for tile_idx in range(TILES_PER_SPLIT):
        kv_offset = split_kv_start + tile_idx * BLOCK_KV
        k_lat_fp8 = tl.load(KV_FP8 + (kv_offset + offs_kv[:, None]) * stride_kv_tok + offs_dk_lat[None, :])
        k_rope_fp8 = tl.load(KV_FP8 + (kv_offset + offs_kv[:, None]) * stride_kv_tok + 512 + offs_dk_rope[None, :])
        scores = (tl.dot(q_lat, tl.trans(k_lat_fp8)) + tl.dot(q_rope, tl.trans(k_rope_fp8))).to(tl.float32) * score_scale
        m_ij = tl.max(scores, axis=1)
        m_new = tl.maximum(m_i, m_ij)
        alpha = tl.math.exp2(m_i - m_new)
        p = tl.math.exp2(scores - m_new[:, None])
        l_i = alpha * l_i + tl.sum(p, axis=1)
        k_lat_bf16 = k_lat_fp8.to(tl.bfloat16)
        acc = acc * alpha[:, None] + tl.dot(p.to(tl.bfloat16), k_lat_bf16)
        m_i = m_new
    acc = acc * kv_s
    tl.store(partial_m + out_base + offs_h, m_i)
    tl.store(partial_l + out_base + offs_h, l_i)
    tl.store(partial_O + (out_base + offs_h[:, None]) * BLOCK_DV + offs_dv[None, :], acc)


@triton.jit
def _flash_decode_stage2(
    partial_O, partial_m, partial_l, O, qo_indptr,
    NUM_SPLITS: tl.constexpr, BLOCK_DV: tl.constexpr,
    NUM_HEADS_CONST: tl.constexpr, BATCH_SIZE: tl.constexpr,
):
    pid = tl.program_id(0)
    batch_id = pid // NUM_HEADS_CONST
    head_id = pid % NUM_HEADS_CONST
    q_idx = tl.load(qo_indptr + batch_id)
    offs_dv = tl.arange(0, BLOCK_DV)
    m_global = tl.full([], value=float("-inf"), dtype=tl.float32)
    for s in range(NUM_SPLITS):
        base = (s * BATCH_SIZE + batch_id) * NUM_HEADS_CONST + head_id
        m_global = tl.maximum(m_global, tl.load(partial_m + base))
    acc = tl.zeros([BLOCK_DV], dtype=tl.float32)
    l_total = tl.full([], value=0.0, dtype=tl.float32)
    for s in range(NUM_SPLITS):
        base = (s * BATCH_SIZE + batch_id) * NUM_HEADS_CONST + head_id
        alpha = tl.math.exp2(tl.load(partial_m + base) - m_global)
        acc += alpha * tl.load(partial_O + base * BLOCK_DV + offs_dv)
        l_total += alpha * tl.load(partial_l + base)
    acc = acc / tl.maximum(l_total, 1e-12)
    tl.store(O + q_idx * NUM_HEADS_CONST * BLOCK_DV + head_id * BLOCK_DV + offs_dv, acc.to(tl.bfloat16))


@triton.jit
def _q_scale_cast_fp8(Q_IN, Q_OUT, inv_scale, N_ELEM: tl.constexpr, BLOCK: tl.constexpr):
    pid = tl.program_id(0)
    offs = pid * BLOCK + tl.arange(0, BLOCK)
    mask = offs < N_ELEM
    q = tl.load(Q_IN + offs, mask=mask).to(tl.float32) * inv_scale
    tl.store(Q_OUT + offs, q.to(tl.float8e4nv), mask=mask)


_triton_buf = {}
_aiter_meta = {}
_aiter_idx = {}
_aiter_kv_last = {}
_aiter_meta_ready = {}
_out_buf = {}
_q_fp8_buf = {}

_finfo = torch.finfo(FP8_DTYPE)
_FIXED_INV_SCALE = _finfo.max / 16.0
_FIXED_Q_SCALE = torch.tensor([16.0 / _finfo.max], dtype=torch.float32, device="cuda")


def _run_fused_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, BKV, kv_seq_len):
    kv_fp8, kv_scale = kv_data["fp8"]
    kv_flat = kv_fp8.view(-1, QK_HEAD_DIM)
    total_q = q.shape[0]
    ok = (total_q, nq)
    if ok not in _out_buf:
        _out_buf[ok] = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
    o = _out_buf[ok]
    _fused_fp8qk_singlepass[(batch_size,)](
        q, kv_flat, kv_scale, o, kv_indptr, qo_indptr, SM_SCALE_LOG2E,
        q.stride(0), q.stride(1), kv_flat.stride(0),
        BLOCK_KV=BKV, BLOCK_DV=V_HEAD_DIM, NUM_HEADS_CONST=nq,
        TILES_TOTAL=kv_seq_len // BKV, num_stages=2, allow_flush_denorm=True)
    return o


def _run_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, NS, BKV, kv_seq_len, stages=2):
    kv_fp8, kv_scale = kv_data["fp8"]
    kv_flat = kv_fp8.view(-1, QK_HEAD_DIM)
    total_q = q.shape[0]
    total_e = NS * batch_size * nq
    ck = ("fp8qk", NS, batch_size, BKV)
    if ck not in _triton_buf:
        _triton_buf[ck] = (
            torch.empty((total_e, V_HEAD_DIM), dtype=torch.float32, device="cuda"),
            torch.empty((total_e,), dtype=torch.float32, device="cuda"),
            torch.empty((total_e,), dtype=torch.float32, device="cuda"),
        )
    pO, pm, pl = _triton_buf[ck]
    tiles_per_split = kv_seq_len // (NS * BKV)
    _flash_decode_fp8qk_stage1[(NS, batch_size)](
        q, kv_flat, kv_scale, pO, pm, pl, kv_indptr, qo_indptr, SM_SCALE_LOG2E,
        q.stride(0), q.stride(1), kv_flat.stride(0),
        NUM_SPLITS=NS, BLOCK_KV=BKV, BLOCK_DV=V_HEAD_DIM, NUM_HEADS_CONST=nq,
        TILES_PER_SPLIT=tiles_per_split, num_stages=stages, allow_flush_denorm=True)
    ok = (total_q, nq)
    if ok not in _out_buf:
        _out_buf[ok] = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
    o = _out_buf[ok]
    _flash_decode_stage2[(batch_size * nq,)](
        pO, pm, pl, o, qo_indptr,
        NUM_SPLITS=NS, BLOCK_DV=V_HEAD_DIM, NUM_HEADS_CONST=nq, BATCH_SIZE=batch_size,
        allow_flush_denorm=True)
    return o


def _run_aiter(q, kv_data, qo_indptr, kv_indptr, config, num_splits=32):
    batch_size = config["batch_size"]
    nq = config["num_heads"]
    nkv = config["num_kv_heads"]
    dq = config["qk_head_dim"]
    dv = config["v_head_dim"]
    q_seq_len = config["q_seq_len"]
    kv_seq_len = config["kv_seq_len"]
    NUM_SPLITS = num_splits
    n_elem = q.numel()
    qk = ("q_fp8", n_elem)
    if qk not in _q_fp8_buf:
        _q_fp8_buf[qk] = torch.empty(q.shape, dtype=FP8_DTYPE, device="cuda")
    q_fp8 = _q_fp8_buf[qk]
    BLOCK = 1024
    grid = ((n_elem + BLOCK - 1) // BLOCK,)
    _q_scale_cast_fp8[grid](q.view(-1), q_fp8.view(-1), _FIXED_INV_SCALE, N_ELEM=n_elem, BLOCK=BLOCK)
    q_scale = _FIXED_Q_SCALE
    kv_fp8, kv_s = kv_data["fp8"]
    total_kv_len = batch_size * kv_seq_len
    kv_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, nkv, kv_fp8.shape[-1])
    kl_key = (batch_size, kv_seq_len)
    if kl_key not in _aiter_kv_last:
        _aiter_kv_last[kl_key] = torch.full(
            (batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")
    kv_last = _aiter_kv_last[kl_key]
    ck = (batch_size, total_kv_len, nq, q_fp8.dtype, kv_fp8.dtype, NUM_SPLITS)
    if ck not in _aiter_meta:
        info = get_mla_metadata_info_v1(
            batch_size, q_seq_len, nq, q_fp8.dtype, kv_fp8.dtype,
            is_sparse=False, fast_mode=False,
            num_kv_splits=NUM_SPLITS, intra_batch_mode=True)
        _aiter_meta[ck] = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
    work = _aiter_meta[ck]
    (wm, wi, wis, ri, rfm, rpm) = work
    if ck not in _aiter_meta_ready:
        get_mla_metadata_v1(
            qo_indptr, kv_indptr, kv_last,
            nq // nkv, nkv, True,
            wm, wis, wi, ri, rfm, rpm,
            page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16),
            max_seqlen_qo=q_seq_len, uni_seqlen_qo=q_seq_len,
            fast_mode=False, max_split_per_batch=NUM_SPLITS,
            intra_batch_mode=True, dtype_q=q_fp8.dtype, dtype_kv=kv_fp8.dtype)
        _aiter_meta_ready[ck] = True
    if total_kv_len not in _aiter_idx:
        _aiter_idx[total_kv_len] = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
    ok = (q.shape[0], nq, dv)
    if ok not in _out_buf:
        _out_buf[ok] = torch.empty(ok, dtype=torch.bfloat16, device="cuda")
    o = _out_buf[ok]
    mla_decode_fwd(
        q_fp8.view(-1, nq, dq), kv_4d, o,
        qo_indptr, kv_indptr, _aiter_idx[total_kv_len],
        kv_last, q_seq_len,
        page_size=PAGE_SIZE, nhead_kv=nkv,
        sm_scale=SM_SCALE, logit_cap=0.0,
        num_kv_splits=NUM_SPLITS,
        q_scale=q_scale, kv_scale=kv_s,
        intra_batch_mode=True,
        work_meta_data=wm, work_indptr=wi, work_info_set=wis,
        reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm)
    return o


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    batch_size = config["batch_size"]
    kv_seq_len = config["kv_seq_len"]
    nq = config["num_heads"]

    if kv_seq_len <= 1024:
        if batch_size <= 4:
            return _run_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, NS=16, BKV=64, kv_seq_len=kv_seq_len, stages=1)
        elif batch_size <= 32:
            return _run_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, NS=8, BKV=64, kv_seq_len=kv_seq_len)
        elif batch_size <= 64:
            return _run_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, NS=4, BKV=64, kv_seq_len=kv_seq_len)
        else:
            return _run_fused_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, BKV=64, kv_seq_len=kv_seq_len)
    else:
        if batch_size <= 4:
            return _run_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, NS=32, BKV=64, kv_seq_len=kv_seq_len)
        elif batch_size <= 32:
            return _run_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, NS=8, BKV=64, kv_seq_len=kv_seq_len)
        elif batch_size <= 64:
            return _run_fp8qk(q, kv_data, qo_indptr, kv_indptr, batch_size, nq, NS=4, BKV=64, kv_seq_len=kv_seq_len)
        else:
            # AITER with num_kv_splits=24 (sweeping down from 32)
            return _run_aiter(q, kv_data, qo_indptr, kv_indptr, config, num_splits=24)
scrolls · 298 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