Skip to content
KernelIndex
Search⌘K

submission 647642

rofeca4922 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

test.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-647642?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
57.8µs
#221 of 766
2026-03-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d67df1580ea4bdefb8162172c837a257a694c6d114c1e4983cb0092c82efdf24
license declaredunknown
license concludedunknown
authorsrofeca4922
imported2026-08-15

Techniques

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

num-warps = 4num_warps=4,
persistent-kernelwork item in the persistent scheduler handles 64 KV tokens instead of 16,
stages = 2num_stages=2,

Kernel source

test.py207 lines
"""
test_v86_kv_granularity: Try kv_granularity=64 for kv=8192 shapes
Base: test.py (v83)
Direction: NEW — kv_granularity parameter exploration
Target: s4/s6/s8 (kv=8192, bandwidth-dominated)
Change: For kv=8192 shapes, use kv_granularity=64 instead of 16. This means each
        work item in the persistent scheduler handles 64 KV tokens instead of 16,
        potentially improving spatial locality and reducing scheduling overhead.
        kv<=1024 shapes keep granularity=16.
Rationale: s8 profile: stage1=237.7us, MfmaUtil=6.5%, OccupancyPerCU=3.6 (bandwidth-bound).
           Larger granularity may improve memory coalescing for large KV sequences.
Scale: MODERATE
"""
import torch
import aiter
import triton
from task import input_t, output_t

from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.mla import get_meta_param, _fwd_kernel_stage2_asm

NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM
V_HEAD_DIM = KV_LORA_RANK
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8

_cache = {}


def _ensure_cache_nonpers_bf16(batch_size, kv_seq_len, total_q):
    key = ("npbf16", batch_size, kv_seq_len)
    if key in _cache:
        return _cache[key]

    nq = NUM_HEADS
    total_kv = batch_size * kv_seq_len

    kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
    num_kv_splits, num_kv_splits_indptr = get_meta_param(None, batch_size, total_kv, nq, 1, torch.bfloat16)
    o = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
    logits = torch.empty((total_q, num_kv_splits, nq, V_HEAD_DIM), dtype=torch.float32, device="cuda")
    attn_lse = torch.empty((total_q, num_kv_splits, nq, 1), dtype=torch.float32, device="cuda")

    _cache[key] = {
        "kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,
        "num_kv_splits": num_kv_splits, "num_kv_splits_indptr": num_kv_splits_indptr,
        "logits": logits, "attn_lse": attn_lse, "o": o,
    }
    return _cache[key]


def _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, persistent_splits, fast_mode, kv_gran=16):
    key = ("pfp8", batch_size, kv_seq_len, persistent_splits, fast_mode, kv_gran)
    if key in _cache:
        return _cache[key]

    max_q_len = 1
    nq, nkv = NUM_HEADS, NUM_KV_HEADS
    total_kv = batch_size * kv_seq_len

    kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")

    info = get_mla_metadata_info_v1(
        batch_size, max_q_len, nq, FP8_DTYPE, FP8_DTYPE,
        is_sparse=False, fast_mode=fast_mode,
        num_kv_splits=persistent_splits, intra_batch_mode=True,
    )
    work = [torch.empty(s, dtype=t, device="cuda") for s, t 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,
        nq // nkv, nkv, 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, kv_gran),
        max_seqlen_qo=max_q_len,
        uni_seqlen_qo=max_q_len,
        fast_mode=fast_mode,
        max_split_per_batch=persistent_splits,
        intra_batch_mode=True,
        dtype_q=FP8_DTYPE,
        dtype_kv=FP8_DTYPE,
    )

    num_partials = reduce_partial_map.size(0)
    logits = torch.empty((num_partials, 1, nq, V_HEAD_DIM), dtype=torch.float32, device="cuda")
    attn_lse = torch.empty((num_partials, 1, nq, 1), dtype=torch.float32, device="cuda")
    o = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
    q_fp8 = torch.empty((total_q, nq * QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda")
    q_scale = torch.ones(1, dtype=torch.float32, device="cuda")

    _cache[key] = {
        "kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,
        "work_metadata": 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,
        "logits": logits, "attn_lse": attn_lse, "o": o,
        "q_fp8": q_fp8, "q_scale": q_scale,
        "num_partials": num_partials,
    }
    return _cache[key]


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"]
    total_q = q.shape[0]
    total_kv = batch_size * kv_seq_len

    # ---- Tier 1: batch<=4 -> bf16/bf16 non-persistent (s1, s2) ----
    if batch_size <= 4:
        kv_bf16 = kv_data["bf16"]
        kv_4d = kv_bf16.view(total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)
        c = _ensure_cache_nonpers_bf16(batch_size, kv_seq_len, total_q)

        aiter.mla_decode_stage1_asm_fwd(
            q.view(-1, NUM_HEADS, QK_HEAD_DIM), kv_4d,
            qo_indptr, kv_indptr, c["kv_indices"], c["kv_last_page_len"],
            c["num_kv_splits_indptr"],
            None, None, None,
            1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
            c["logits"], c["attn_lse"], c["o"],
            None, None,
        )

        Lv = V_HEAD_DIM
        BLOCK_DV = triton.next_power_of_2(Lv)
        _fwd_kernel_stage2_asm[(batch_size, NUM_HEADS)](
            c["logits"], c["attn_lse"], c["o"],
            qo_indptr, kv_indptr, c["num_kv_splits_indptr"],
            c["attn_lse"].stride(0), c["attn_lse"].stride(2), c["attn_lse"].stride(1),
            c["o"].stride(0), c["o"].stride(1),
            MAYBE_FINAL_OUT=True,
            BATCH_NUM=batch_size,
            BLOCK_DV=BLOCK_DV,
            Lv=Lv,
            mgc=64,
            num_warps=4,
            num_stages=2,
            waves_per_eu=4,
        )
        return c["o"]

    # ---- Tier 2: ALL fp8 shapes -> persistent (s3-s8) ----
    # Non-persistent fp8 was faster but fails leaderboard correctness (v77, v78).
    # Persistent + mla_reduce_v1 is the only leaderboard-safe fp8 path.
    else:
        kv_buffer_fp8, kv_scale = kv_data["fp8"]
        kv_buffer_4d = kv_buffer_fp8.view(total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)

        # Per-shape split tuning
        if total_kv >= 1000000:
            # s8 (256, 8192) -- splits=4 (from v77)
            splits, fast_mode = 4, False
        elif total_kv >= 300000:
            # s6 (64, 8192)
            splits, fast_mode = 8, False
        elif batch_size >= 256:
            # s7 (256, 1024) -- splits=1 fills 256 CUs directly
            splits, fast_mode = 1, False
        elif batch_size >= 64:
            # s5 (64, 1024) -- splits=4 (reverted from splits=1 which regressed in v82)
            splits, fast_mode = 4, False
        else:
            # s3 (32, 1024) and s4 (32, 8192)
            if kv_seq_len <= 1024:
                splits, fast_mode = 4, True   # s3: reduced from 8 to 4
            else:
                splits, fast_mode = 32, True  # s4

        # Use larger kv_granularity for long KV sequences (64 instead of 16)
        kv_gran = 64 if kv_seq_len > 1024 else 16
        c = _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, splits, fast_mode, kv_gran)

        # Fast FP8 quant: copy_ cast (scale=1.0) -- from v63
        q_2d = q.view(total_q, NUM_HEADS * QK_HEAD_DIM)
        c["q_fp8"].copy_(q_2d)

        aiter.mla_decode_stage1_asm_fwd(
            c["q_fp8"].view(-1, NUM_HEADS, QK_HEAD_DIM), kv_buffer_4d,
            qo_indptr, kv_indptr, c["kv_indices"], c["kv_last_page_len"],
            None, c["work_metadata"], c["work_indptr"], c["work_info_set"],
            1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
            c["logits"], c["attn_lse"], c["o"],
            c["q_scale"], kv_scale,
        )

        aiter.mla_reduce_v1(
            c["logits"], c["attn_lse"],
            c["reduce_indptr"], c["reduce_final_map"], c["reduce_partial_map"],
            1, c["o"], None,
        )
        return c["o"]
scrolls · 207 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