Skip to content
KernelIndex
Search⌘K

submission 747989

bigmodel_wuzhigang · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

team_mla_v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-747989?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
35.4µs
#77 of 766
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3f1a719e35761cd47176a7e1d5633e9af23d9231ffb5e2487a5eb514ba7f2eac
license declaredunknown
license concludedunknown
authorsbigmodel_wuzhigang
imported2026-08-15

Kernel source

team_mla_v3.py280 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""
Team MLA trunk v3.

Single-knob exact-family A/B on top of team_mla_v1:
- keep the same `256/1024` pg2 + bf16Q route
- change only `num_kv_splits` for that route from 4 to 8
- keep all other shapes and policies unchanged
"""

import torch
import triton
import triton.language as tl

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

from task import input_t, output_t

FP8_DTYPE = aiter_dtypes.fp8
BF16 = torch.bfloat16
FP8_MAX = float(torch.finfo(FP8_DTYPE).max)
FIXED_AMAX = 16.0
PG2_256_NUM_KV_SPLITS = 8
PG2_256_KV_GRANULARITY = 8

_meta_cache = {}
_alloc_cache = {}
_bf16_np_cache = {}


@triton.jit
def _q_to_fp8_kernel(
    q_ptr,
    out_ptr,
    scale_ptr,
    amax_ptr,
    FP8_MAX: tl.constexpr,
    N,
    BLOCK: tl.constexpr,
):
    amax = tl.load(amax_ptr)
    amax = tl.where(amax < 1e-12, 1e-12, amax)
    scale = amax / FP8_MAX
    if tl.program_id(0) == 0:
        tl.store(scale_ptr, scale)
    pid = tl.program_id(0)
    offs = pid * BLOCK + tl.arange(0, BLOCK)
    mask = offs < N
    x = tl.load(q_ptr + offs, mask=mask, other=0.0).to(tl.float32)
    x = x / scale
    x = tl.clamp(x, -FP8_MAX, FP8_MAX)
    tl.store(out_ptr + offs, x.to(out_ptr.dtype.element_ty), mask=mask)


def _build_meta(
    batch_size,
    kv_seq_len,
    q_seq_len,
    nq,
    nkv,
    num_kv_splits,
    page_size,
    dtype_q,
    dtype_kv,
    qo_indptr,
    kv_indptr,
    kv_granularity,
):
    total_kv = batch_size * kv_seq_len
    if page_size == 1:
        num_pages = total_kv
        kv_indptr_pages = kv_indptr
        seq_lens = kv_indptr[1:] - kv_indptr[:-1]
        kv_last_page_len = seq_lens.to(torch.int32)
    else:
        num_pages = total_kv // page_size
        kv_indptr_pages = kv_indptr // page_size
        seq_lens = kv_indptr[1:] - kv_indptr[:-1]
        kv_last_page_len = (seq_lens % page_size).to(torch.int32)
        kv_last_page_len = torch.where(kv_last_page_len == 0, page_size, kv_last_page_len)

    info = get_mla_metadata_info_v1(
        batch_size,
        q_seq_len,
        nq,
        dtype_q,
        dtype_kv,
        is_sparse=False,
        fast_mode=False,
        num_kv_splits=num_kv_splits,
        intra_batch_mode=True,
    )
    work = [torch.empty(shape, dtype=dtype, device="cuda") for shape, dtype in info]
    wm, wi, wis, ri, rfm, rpm = work
    get_mla_metadata_v1(
        qo_indptr,
        kv_indptr_pages,
        kv_last_page_len,
        nq // nkv,
        nkv,
        True,
        wm,
        wis,
        wi,
        ri,
        rfm,
        rpm,
        page_size=page_size,
        kv_granularity=kv_granularity,
        max_seqlen_qo=q_seq_len,
        uni_seqlen_qo=q_seq_len,
        fast_mode=False,
        max_split_per_batch=num_kv_splits,
        intra_batch_mode=True,
        dtype_q=dtype_q,
        dtype_kv=dtype_kv,
    )
    kv_indices = torch.arange(num_pages, dtype=torch.int32, device="cuda")
    return wm, wi, wis, ri, rfm, rpm, kv_indices, kv_last_page_len, kv_indptr_pages, page_size


def _choose_num_kv_splits(batch_size: int, kv_seq_len: int) -> int:
    if batch_size == 256 and kv_seq_len == 1024:
        return PG2_256_NUM_KV_SPLITS
    if batch_size <= 4:
        return 4
    if batch_size <= 64:
        return 8
    return 16


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    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"]
    sm_scale = config["sm_scale"]
    kv_seq_len = config["kv_seq_len"]

    use_small_bf16_np = batch_size <= 32 and kv_seq_len <= 1024
    use_bf16_pg2_256 = batch_size == 256 and kv_seq_len == 1024

    if use_small_bf16_np:
        cache_key = (batch_size, kv_seq_len)
        if cache_key not in _bf16_np_cache:
            total_kv = batch_size * kv_seq_len
            kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
            kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")
            _bf16_np_cache[cache_key] = (kv_indices, kv_last_page_len)
        kv_indices, kv_last_page_len = _bf16_np_cache[cache_key]
        alloc_key = ("bf16_np", q.shape[0], nq, dv)
        if alloc_key not in _alloc_cache:
            _alloc_cache[alloc_key] = torch.empty((q.shape[0], nq, dv), dtype=BF16, device="cuda")
        o = _alloc_cache[alloc_key]
        kv_buffer_bf16 = kv_data["bf16"].view(-1, 1, nkv, dq)
        mla_decode_fwd(
            q,
            kv_buffer_bf16,
            o,
            qo_indptr,
            kv_indptr,
            kv_indices,
            kv_last_page_len,
            q_seq_len,
            page_size=1,
            nhead_kv=nkv,
            sm_scale=sm_scale,
            intra_batch_mode=False,
        )
        return o

    if use_bf16_pg2_256:
        page_size = 2
        dtype_q = BF16
        dtype_kv = FP8_DTYPE
        kv_granularity = PG2_256_KV_GRANULARITY
        use_fp8_q = False
    else:
        page_size = 1 if kv_seq_len <= 1024 else 8
        dtype_q = FP8_DTYPE
        dtype_kv = FP8_DTYPE
        kv_granularity = max(1, 16 // page_size)
        use_fp8_q = True

    num_kv_splits = _choose_num_kv_splits(batch_size, kv_seq_len)

    cache_key = (
        batch_size,
        kv_seq_len,
        num_kv_splits,
        page_size,
        str(dtype_q),
        str(dtype_kv),
        kv_granularity,
    )
    if cache_key not in _meta_cache:
        _meta_cache[cache_key] = _build_meta(
            batch_size,
            kv_seq_len,
            q_seq_len,
            nq,
            nkv,
            num_kv_splits,
            page_size,
            dtype_q,
            dtype_kv,
            qo_indptr,
            kv_indptr,
            kv_granularity,
        )

    wm, wi, wis, ri, rfm, rpm, kv_indices, kv_last_page_len, kv_indptr_pages, ps = _meta_cache[cache_key]
    kv_buffer_fp8, kv_scale = kv_data["fp8"]
    kv_buffer_4d = kv_buffer_fp8.view(-1, ps, nkv, kv_buffer_fp8.shape[-1])

    if use_fp8_q:
        alloc_key = ("fp8_structmix", q.shape[0], nq, dv, dq)
        if alloc_key not in _alloc_cache:
            amax_buf = torch.full((1,), FIXED_AMAX, dtype=torch.float32, device="cuda")
            _alloc_cache[alloc_key] = (
                torch.empty((q.shape[0], nq, dv), dtype=BF16, device="cuda"),
                amax_buf,
                torch.empty(1, dtype=torch.float32, device="cuda"),
                torch.empty(q.shape[0] * nq * dq, dtype=FP8_DTYPE, device="cuda"),
            )
        o, amax_buf, scale_buf, q_fp8_flat = _alloc_cache[alloc_key]
        n = q.numel()
        block = 4096
        grid = ((n + block - 1) // block,)
        _q_to_fp8_kernel[grid](
            q,
            q_fp8_flat,
            scale_buf,
            amax_buf,
            FP8_MAX=FP8_MAX,
            N=n,
            BLOCK=block,
        )
        q_input = q_fp8_flat.view(q.shape[0], nq, dq)
        kwargs = {"q_scale": scale_buf, "kv_scale": kv_scale}
    else:
        alloc_key = ("bf16_pg2split8_256", q.shape[0], nq, dv)
        if alloc_key not in _alloc_cache:
            _alloc_cache[alloc_key] = torch.empty((q.shape[0], nq, dv), dtype=BF16, device="cuda")
        o = _alloc_cache[alloc_key]
        q_input = q
        kwargs = {"kv_scale": kv_scale}

    mla_decode_fwd(
        q_input,
        kv_buffer_4d,
        o,
        qo_indptr,
        kv_indptr_pages,
        kv_indices,
        kv_last_page_len,
        q_seq_len,
        page_size=ps,
        nhead_kv=nkv,
        sm_scale=sm_scale,
        logit_cap=0.0,
        num_kv_splits=num_kv_splits,
        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,
        **kwargs,
    )
    return o
scrolls · 280 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