Skip to content
KernelIndex
Search⌘K

submission 653808

lonk · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

amd-mixed-mla-hybrid-v57-legacy28.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-653808?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
64.4µs
#280 of 766
2026-03-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:be044a8b768ac8aada83434b39f25856d12eb9a2efd71ff847d2c62f18164022
license declaredunknown
license concludedunknown
authorslonk
imported2026-08-26

Techniques

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

mmascores = tl.dot(q_lat, tl.trans(kv_lat))
num-warps = 1num_warps=1,
persistent-kerneldef _persistent_decode(data: input_t) -> output_t:

Kernel source

amd-mixed-mla-hybrid-v57-legacy28.py392 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t

# Hybrid v57 legacy-28 sweep: keep the validated Triton baseline for
# seven shapes and route the gating shape through the legacy 28-split metadata regime
# without any cross-call caches.

PAGE_SIZE = 1
MAX_KV_SPLITS = 32
TARGET_TILE_COUNT = 512
MIN_KV_TOKENS_PER_SPLIT = 256


def _aiter_symbols():
    from aiter import dtypes as aiter_dtypes
    from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
    from aiter.mla import mla_decode_fwd

    return aiter_dtypes.fp8, get_mla_metadata_info_v1, get_mla_metadata_v1, mla_decode_fwd


def _quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    fp8_dtype, _, _, _ = _aiter_symbols()
    finfo = torch.finfo(fp8_dtype)
    amax = tensor.abs().amax().clamp(min=1e-12)
    scale = (amax / finfo.max).to(torch.float32).reshape(1)
    fp8_tensor = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(fp8_dtype)
    return fp8_tensor, scale


def _choose_num_kv_splits(batch_size: int, kv_seq_len: int) -> int:
    del batch_size, kv_seq_len
    return 28


def _make_metadata(
    batch_size: int,
    max_q_len: int,
    nhead: int,
    nhead_kv: int,
    q_dtype: torch.dtype,
    kv_dtype: torch.dtype,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    kv_last_page_len: torch.Tensor,
    num_kv_splits: int,
) -> dict[str, torch.Tensor]:
    _, get_mla_metadata_info_v1, get_mla_metadata_v1, _ = _aiter_symbols()

    info = get_mla_metadata_info_v1(
        batch_size,
        max_q_len,
        nhead,
        q_dtype,
        kv_dtype,
        is_sparse=False,
        fast_mode=False,
        num_kv_splits=num_kv_splits,
        intra_batch_mode=True,
    )
    work = [torch.empty(shape, dtype=dtype, device=qo_indptr.device) for shape, dtype in info]
    (
        work_metadata,
        work_indptr,
        work_info_set,
        reduce_indptr,
        reduce_final_map,
        reduce_partial_map,
    ) = work

    get_mla_metadata_v1(
        qo_indptr,
        kv_indptr,
        kv_last_page_len,
        nhead // nhead_kv,
        nhead_kv,
        True,
        work_metadata,
        work_info_set,
        work_indptr,
        reduce_indptr,
        reduce_final_map,
        reduce_partial_map,
        page_size=PAGE_SIZE,
        kv_granularity=max(PAGE_SIZE, 16),
        max_seqlen_qo=max_q_len,
        uni_seqlen_qo=max_q_len,
        fast_mode=False,
        max_split_per_batch=num_kv_splits,
        intra_batch_mode=True,
        dtype_q=q_dtype,
        dtype_kv=kv_dtype,
    )

    return {
        "work_meta_data": work_metadata,
        "work_indptr": work_indptr,
        "work_info_set": work_info_set,
        "reduce_indptr": reduce_indptr,
        "reduce_final_map": reduce_final_map,
        "reduce_partial_map": reduce_partial_map,
    }


def _persistent_decode(data: input_t) -> output_t:
    _, _, _, mla_decode_fwd = _aiter_symbols()

    q, kv_data, qo_indptr, kv_indptr, config = data
    q_fp8, q_scale = _quantize_fp8(q)
    kv_buffer_fp8, kv_scale = kv_data["fp8"]

    batch_size = config["batch_size"]
    num_heads = config["num_heads"]
    num_kv_heads = config["num_kv_heads"]
    q_seq_len = config["q_seq_len"]
    kv_seq_len = config["kv_seq_len"]
    total_kv = kv_buffer_fp8.shape[0]
    num_kv_splits = _choose_num_kv_splits(batch_size, kv_seq_len)

    kv_indices = torch.arange(total_kv, dtype=torch.int32, device=q.device)
    kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
    out = torch.empty((q.shape[0], num_heads, config["v_head_dim"]), dtype=torch.bfloat16, device=q.device)
    meta = _make_metadata(
        batch_size,
        q_seq_len,
        num_heads,
        num_kv_heads,
        q_fp8.dtype,
        kv_buffer_fp8.dtype,
        qo_indptr,
        kv_indptr,
        kv_last_page_len,
        num_kv_splits,
    )

    mla_decode_fwd(
        q_fp8,
        kv_buffer_fp8.view(total_kv, PAGE_SIZE, num_kv_heads, kv_buffer_fp8.shape[-1]),
        out,
        qo_indptr,
        kv_indptr,
        kv_indices,
        kv_last_page_len,
        q_seq_len,
        page_size=PAGE_SIZE,
        nhead_kv=num_kv_heads,
        sm_scale=config["sm_scale"],
        logit_cap=0.0,
        num_kv_splits=num_kv_splits,
        q_scale=q_scale,
        kv_scale=kv_scale,
        intra_batch_mode=True,
        **meta,
    )
    return out


@triton.jit
def _mla_s1(
    Q,
    KV,
    KV_SCALE,
    QO_ind,
    KV_ind,
    PO,
    PM,
    PL,
    sm_scale,
    sq0,
    sq1,
    skv,
    spo0,
    spo1,
    spo2,
    sml0,
    sml1,
    NS: tl.constexpr,
    BK: tl.constexpr,
):
    bid = tl.program_id(0)
    sid = tl.program_id(1)

    q_off = tl.load(QO_ind + bid)
    kv_lo = tl.load(KV_ind + bid)
    kv_hi = tl.load(KV_ind + bid + 1)
    kv_n = kv_hi - kv_lo

    kvs = tl.load(KV_SCALE)
    csz = tl.cdiv(kv_n, NS)
    lo = kv_lo + sid * csz
    hi = tl.minimum(lo + csz, kv_hi)
    cs = kvs * sm_scale

    h16 = tl.arange(0, 16)
    d512 = tl.arange(0, 512)
    d64 = tl.arange(0, 64)

    q_lat = tl.load(Q + q_off * sq0 + h16[:, None] * sq1 + d512[None, :])
    q_rope = tl.load(Q + q_off * sq0 + h16[:, None] * sq1 + 512 + d64[None, :])

    mi = tl.full([16], float("-inf"), dtype=tl.float32)
    li = tl.zeros([16], dtype=tl.float32)
    acc = tl.zeros([16, 512], dtype=tl.float32)

    bk_range = tl.arange(0, BK)

    pos = lo
    while pos < hi:
        remaining = hi - pos
        mask = bk_range < remaining

        kv_offs = pos + bk_range
        kv_lat = tl.load(
            KV + kv_offs[:, None] * skv + d512[None, :],
            mask=mask[:, None],
            other=0.0,
        ).to(tl.bfloat16)
        kv_rope = tl.load(
            KV + kv_offs[:, None] * skv + 512 + d64[None, :],
            mask=mask[:, None],
            other=0.0,
        ).to(tl.bfloat16)

        scores = tl.dot(q_lat, tl.trans(kv_lat))
        scores += tl.dot(q_rope, tl.trans(kv_rope))
        scores *= cs
        scores = tl.where(mask[None, :], scores, float("-inf"))

        block_max = tl.max(scores, axis=1)
        mn = tl.maximum(mi, block_max)
        alpha = tl.exp(mi - mn)
        exp_s = tl.exp(scores - mn[:, None])
        exp_s = tl.where(mask[None, :], exp_s, 0.0)

        li = li * alpha + tl.sum(exp_s, axis=1)
        acc = acc * alpha[:, None]
        acc += tl.dot(exp_s.to(tl.bfloat16), kv_lat)

        mi = mn
        pos += BK

    po_result = (acc * kvs).to(tl.bfloat16)
    tl.store(PO + q_off * spo0 + sid * spo1 + h16[:, None] * spo2 + d512[None, :], po_result)
    tl.store(PM + q_off * sml0 + sid * sml1 + h16, mi)
    tl.store(PL + q_off * sml0 + sid * sml1 + h16, li)


@triton.jit
def _mla_s2(
    PO,
    PM,
    PL,
    Out,
    QO_ind,
    spo0,
    spo1,
    spo2,
    sml0,
    sml1,
    so0,
    so1,
    NS: tl.constexpr,
):
    hid = tl.program_id(0)
    bid = tl.program_id(1)

    q_off = tl.load(QO_ind + bid)
    mlb = q_off * sml0 + hid

    gm = float("-inf")
    for s in range(NS):
        gm = tl.maximum(gm, tl.load(PM + mlb + s * sml1))

    acc = tl.zeros([512], dtype=tl.float32)
    lsum = 0.0
    for s in range(NS):
        ms = tl.load(PM + mlb + s * sml1)
        ls = tl.load(PL + mlb + s * sml1)
        w = tl.exp(ms - gm)
        off = q_off * spo0 + s * spo1 + hid * spo2
        po = tl.load(PO + off + tl.arange(0, 512)).to(tl.float32)
        acc += w * po
        lsum += w * ls

    acc = acc / lsum
    off_o = q_off * so0 + hid * so1
    tl.store(Out + off_o + tl.arange(0, 512), acc.to(tl.bfloat16))


def _triton_decode(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    kv_fp8, kv_scale = kv_data["fp8"]

    bs = config["batch_size"]
    nh = config["num_heads"]
    vd = config["v_head_dim"]
    qkd = config["qk_head_dim"]
    sm = config["sm_scale"]
    tq = q.shape[0]
    total_kv = kv_fp8.shape[0]
    kv_seq_len = total_kv // bs

    out = torch.empty((tq, nh, vd), dtype=torch.bfloat16, device=q.device)

    if bs >= 128 and kv_seq_len >= 4096:
        BK = 32
        ns = 8
        nw = 4
    elif bs >= 128:
        BK = 32
        ns = 2
        nw = 2
    elif bs == 64 and kv_seq_len >= 4096:
        BK = 32
        ns = 8
        nw = 2
    elif bs == 64:
        BK = 16
        ns = 8
        nw = 4
    elif bs == 32 and kv_seq_len >= 4096:
        BK = 32
        ns = 16
        nw = 2
    elif bs == 32:
        BK = 16
        ns = 8
        nw = 4
    elif bs <= 8 and kv_seq_len >= 4096:
        BK = 32
        ns = 32 if bs <= 4 else 16
        nw = 4
    else:
        BK = 16
        ns = 16 if bs <= 16 else 8
        nw = 4

    po = torch.empty((tq, ns, nh, vd), dtype=torch.bfloat16, device=q.device)
    pm = torch.empty((tq, ns, nh), dtype=torch.float32, device=q.device)
    pl = torch.empty((tq, ns, nh), dtype=torch.float32, device=q.device)

    _mla_s1[(bs, ns)](
        q,
        kv_fp8,
        kv_scale,
        qo_indptr,
        kv_indptr,
        po,
        pm,
        pl,
        sm,
        q.stride(0),
        q.stride(1),
        qkd,
        po.stride(0),
        po.stride(1),
        po.stride(2),
        pm.stride(0),
        pm.stride(1),
        NS=ns,
        BK=BK,
        num_warps=nw,
    )

    _mla_s2[(nh, bs)](
        po,
        pm,
        pl,
        out,
        qo_indptr,
        po.stride(0),
        po.stride(1),
        po.stride(2),
        pm.stride(0),
        pm.stride(1),
        out.stride(0),
        out.stride(1),
        NS=ns,
        num_warps=1,
    )

    return out


def custom_kernel(data: input_t) -> output_t:
    _, _, _, _, config = data
    if config["batch_size"] == 256 and config["kv_seq_len"] == 8192:
        return _persistent_decode(data)
    return _triton_decode(data)
scrolls · 392 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