Skip to content
KernelIndex
Search⌘K

submission 718377

jotod92140 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

_bss_merged_s4_test_asm_s4_v102_gran128.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-718377?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.1µs
#209 of 766
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:6bd9ca38f064137c43707e5f45845d302d5aeba8c5cd1d98e843e2140d45002b
license declaredunknown
license concludedunknown
authorsjotod92140
imported2026-08-15

Techniques

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

num-warps = 4…K_DV,\n Lv=Lv,\n mgc=64,\n num_warps=4,\n num_stages=2,\n waves_per_eu=4,\n )\n return c["o"]\n\n # ---- Tie…
persistent-kernel…PE = aiter_dtypes.fp8\n\n_cache = {}\n\n\ndef _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, persistent_splits, fast_mode, kv_gran):\n key =…
stages = 2…n mgc=64,\n num_warps=4,\n num_stages=2,\n waves_per_eu=4,\n )\n return c["o"]\n\n # ---- Tier 2: ALL fp8 shapes -> per…

Kernel source

_bss_merged_s4_test_asm_s4_v102_gran128.py41 lines
# Auto-generated by submit-single-shape.py
# Target shape: s4 = {"batchsize": 32, "kvseqlen": 8192, "qseqlen": 1}

import importlib.util
import sys
import os
from pathlib import Path
import tempfile

_TARGET_SOURCE = '"""\ntest_asm_s4_v102_gran128: kv_granularity 64→128\nBase: test_asm_s4_v94_nofm.py\nDirection: NEW — coarser kv granularity\nTarget: s4 (batch=32, kv_seq_len=8192)\nChange: kv_gran 64→128. With 8192 tokens and 32 splits, kv_gran=128\n        gives 64 blocks total, 2 per split — perfectly aligned.\n        kv_gran=64 gives 128 blocks, 4 per split.\n        Coarser granularity reduces metadata overhead.\nRationale: kv_gran=32 regressed +8.7% (iter 548). kv_gran=128 on s6\n           regressed +2.5% (iter 534). s4 kv_gran=128 NEVER tested.\nScale: INCREMENTAL\n"""\nimport torch\nimport aiter\nimport triton\nfrom task import input_t, output_t\n\nfrom aiter import dtypes as aiter_dtypes\nfrom aiter import get_mla_metadata_info_v1, get_mla_metadata_v1\n\nNUM_HEADS = 16\nNUM_KV_HEADS = 1\nKV_LORA_RANK = 512\nQK_ROPE_HEAD_DIM = 64\nQK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM\nV_HEAD_DIM = KV_LORA_RANK\nSM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)\nPAGE_SIZE = 1\nFP8_DTYPE = aiter_dtypes.fp8\n\n_cache = {}\n\n\ndef _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, persistent_splits, fast_mode, kv_gran):\n    key = ("s4_v102_gran128", batch_size, kv_seq_len, persistent_splits, fast_mode, kv_gran)\n    if key in _cache:\n        return _cache[key]\n\n    max_q_len = 1\n    nq, nkv = NUM_HEADS, NUM_KV_HEADS\n    total_kv = batch_size * kv_seq_len\n\n    kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")\n    kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")\n\n    info = get_mla_metadata_info_v1(\n        batch_size, max_q_len, nq, FP8_DTYPE, FP8_DTYPE,\n        is_sparse=False, fast_mode=fast_mode,\n        num_kv_splits=persistent_splits, intra_batch_mode=True,\n    )\n    work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]\n    (work_metadata, work_indptr, work_info_set,\n     reduce_indptr, reduce_final_map, reduce_partial_map) = work\n\n    get_mla_metadata_v1(\n        qo_indptr, kv_indptr, kv_last_page_len,\n        nq // nkv, nkv, True,\n        work_metadata, work_info_set, work_indptr,\n        reduce_indptr, reduce_final_map, reduce_partial_map,\n        page_size=PAGE_SIZE,\n        kv_granularity=max(PAGE_SIZE, kv_gran),\n        max_seqlen_qo=max_q_len,\n        uni_seqlen_qo=max_q_len,\n        fast_mode=fast_mode,\n        max_split_per_batch=persistent_splits,\n        intra_batch_mode=True,\n        dtype_q=FP8_DTYPE,\n        dtype_kv=FP8_DTYPE,\n    )\n\n    num_partials = reduce_partial_map.size(0)\n    logits = torch.empty((num_partials, 1, nq, V_HEAD_DIM), dtype=torch.float32, device="cuda")\n    attn_lse = torch.empty((num_partials, 1, nq, 1), dtype=torch.float32, device="cuda")\n    o = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")\n    q_fp8 = torch.empty((total_q, nq * QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda")\n    q_scale = torch.ones(1, dtype=torch.float32, device="cuda")\n\n    _cache[key] = {\n        "kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,\n        "work_metadata": work_metadata, "work_indptr": work_indptr,\n        "work_info_set": work_info_set, "reduce_indptr": reduce_indptr,\n        "reduce_final_map": reduce_final_map, "reduce_partial_map": reduce_partial_map,\n        "logits": logits, "attn_lse": attn_lse, "o": o,\n        "q_fp8": q_fp8, "q_scale": q_scale,\n        "num_partials": num_partials,\n    }\n    return _cache[key]\n\n\ndef custom_kernel(data: input_t) -> output_t:\n    q, kv_data, qo_indptr, kv_indptr, config = data\n\n    batch_size = config["batch_size"]\n    kv_seq_len = config["kv_seq_len"]\n    total_q = q.shape[0]\n\n    if not (batch_size == 32 and kv_seq_len == 8192):\n        return torch.zeros(\n            (total_q, NUM_HEADS, V_HEAD_DIM),\n            dtype=torch.bfloat16,\n            device=q.device,\n        )\n\n    total_kv = batch_size * kv_seq_len\n\n    splits = 32\n    fast_mode = False\n    kv_gran = 128       # KEY CHANGE: was 64\n\n    kv_buffer_fp8, kv_scale = kv_data["fp8"]\n    kv_buffer_4d = kv_buffer_fp8.view(total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)\n\n    c = _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, splits, fast_mode, kv_gran)\n\n    q_2d = q.view(total_q, NUM_HEADS * QK_HEAD_DIM)\n    c["q_fp8"].copy_(q_2d)\n\n    aiter.mla_decode_stage1_asm_fwd(\n        c["q_fp8"].view(-1, NUM_HEADS, QK_HEAD_DIM), kv_buffer_4d,\n        qo_indptr, kv_indptr, c["kv_indices"], c["kv_last_page_len"],\n        None, c["work_metadata"], c["work_indptr"], c["work_info_set"],\n        1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,\n        c["logits"], c["attn_lse"], c["o"],\n        c["q_scale"], kv_scale,\n    )\n\n    aiter.mla_reduce_v1(\n        c["logits"], c["attn_lse"],\n        c["reduce_indptr"], c["reduce_final_map"], c["reduce_partial_map"],\n        1, c["o"], None,\n    )\n    return c["o"]\n'
_REF_SOURCE = '"""\ntest_v143_s6_splits4: s6 splits 8→4 (continue reduce optimization pattern)\nBase: test.py (v142)\nDirection: NEW — s6 splits tuning\nTarget: s6 (64,8192) — reduce overhead with kv=8192\nChange: s6 splits 8→4. batch=64 × splits=4 = 256 programs (100% CU fill).\n        Follows s5 splits reduction pattern (v142 +8.5%).\nRationale: v140 profile s8 reduce=3.3us. s6 with splits=8 has more reduce overhead.\nScale: INCREMENTAL\n"""\nimport torch\nimport aiter\nimport triton\nfrom task import input_t, output_t\n\nfrom aiter import dtypes as aiter_dtypes\nfrom aiter import get_mla_metadata_info_v1, get_mla_metadata_v1\nfrom aiter.mla import get_meta_param, _fwd_kernel_stage2_asm\n\nNUM_HEADS = 16\nNUM_KV_HEADS = 1\nKV_LORA_RANK = 512\nQK_ROPE_HEAD_DIM = 64\nQK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM\nV_HEAD_DIM = KV_LORA_RANK\nSM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)\nPAGE_SIZE = 1\nFP8_DTYPE = aiter_dtypes.fp8\n\n_cache = {}\n\n\ndef _ensure_cache_nonpers_bf16(batch_size, kv_seq_len, total_q):\n    key = ("npbf16", batch_size, kv_seq_len)\n    if key in _cache:\n        return _cache[key]\n\n    nq = NUM_HEADS\n    total_kv = batch_size * kv_seq_len\n\n    kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")\n    kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")\n    num_kv_splits, num_kv_splits_indptr = get_meta_param(None, batch_size, total_kv, nq, 1, torch.bfloat16)\n    o = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")\n    logits = torch.empty((total_q, num_kv_splits, nq, V_HEAD_DIM), dtype=torch.float32, device="cuda")\n    attn_lse = torch.empty((total_q, num_kv_splits, nq, 1), dtype=torch.float32, device="cuda")\n\n    _cache[key] = {\n        "kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,\n        "num_kv_splits": num_kv_splits, "num_kv_splits_indptr": num_kv_splits_indptr,\n        "logits": logits, "attn_lse": attn_lse, "o": o,\n    }\n    return _cache[key]\n\n\ndef _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, persistent_splits, fast_mode, kv_gran=16):\n    key = ("pfp8", batch_size, kv_seq_len, persistent_splits, fast_mode, kv_gran)\n    if key in _cache:\n        return _cache[key]\n\n    max_q_len = 1\n    nq, nkv = NUM_HEADS, NUM_KV_HEADS\n    total_kv = batch_size * kv_seq_len\n\n    kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")\n    kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")\n\n    info = get_mla_metadata_info_v1(\n        batch_size, max_q_len, nq, FP8_DTYPE, FP8_DTYPE,\n        is_sparse=False, fast_mode=fast_mode,\n        num_kv_splits=persistent_splits, intra_batch_mode=True,\n    )\n    work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]\n    (work_metadata, work_indptr, work_info_set,\n     reduce_indptr, reduce_final_map, reduce_partial_map) = work\n\n    get_mla_metadata_v1(\n        qo_indptr, kv_indptr, kv_last_page_len,\n        nq // nkv, nkv, True,\n        work_metadata, work_info_set, work_indptr,\n        reduce_indptr, reduce_final_map, reduce_partial_map,\n        page_size=PAGE_SIZE,\n        kv_granularity=max(PAGE_SIZE, kv_gran),\n        max_seqlen_qo=max_q_len,\n        uni_seqlen_qo=max_q_len,\n        fast_mode=fast_mode,\n        max_split_per_batch=persistent_splits,\n        intra_batch_mode=True,\n        dtype_q=FP8_DTYPE,\n        dtype_kv=FP8_DTYPE,\n    )\n\n    num_partials = reduce_partial_map.size(0)\n    logits = torch.empty((num_partials, 1, nq, V_HEAD_DIM), dtype=torch.float32, device="cuda")\n    attn_lse = torch.empty((num_partials, 1, nq, 1), dtype=torch.float32, device="cuda")\n    o = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")\n    q_fp8 = torch.empty((total_q, nq * QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda")\n    q_scale = torch.ones(1, dtype=torch.float32, device="cuda")\n\n    _cache[key] = {\n        "kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,\n        "work_metadata": work_metadata, "work_indptr": work_indptr,\n        "work_info_set": work_info_set, "reduce_indptr": reduce_indptr,\n        "reduce_final_map": reduce_final_map, "reduce_partial_map": reduce_partial_map,\n        "logits": logits, "attn_lse": attn_lse, "o": o,\n        "q_fp8": q_fp8, "q_scale": q_scale,\n        "num_partials": num_partials,\n    }\n    return _cache[key]\n\n\ndef custom_kernel(data: input_t) -> output_t:\n    q, kv_data, qo_indptr, kv_indptr, config = data\n\n    batch_size = config["batch_size"]\n    kv_seq_len = config["kv_seq_len"]\n    total_q = q.shape[0]\n    total_kv = batch_size * kv_seq_len\n\n    # ---- Tier 1: batch<=4 -> bf16/bf16 non-persistent (s1, s2) ----\n    if batch_size <= 4:\n        kv_bf16 = kv_data["bf16"]\n        kv_4d = kv_bf16.view(total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)\n        c = _ensure_cache_nonpers_bf16(batch_size, kv_seq_len, total_q)\n\n        aiter.mla_decode_stage1_asm_fwd(\n            q.view(-1, NUM_HEADS, QK_HEAD_DIM), kv_4d,\n            qo_indptr, kv_indptr, c["kv_indices"], c["kv_last_page_len"],\n            c["num_kv_splits_indptr"],\n            None, None, None,\n            1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,\n            c["logits"], c["attn_lse"], c["o"],\n            None, None,\n        )\n\n        Lv = V_HEAD_DIM\n        BLOCK_DV = triton.next_power_of_2(Lv)\n        _fwd_kernel_stage2_asm[(batch_size, NUM_HEADS)](\n            c["logits"], c["attn_lse"], c["o"],\n            qo_indptr, kv_indptr, c["num_kv_splits_indptr"],\n            c["attn_lse"].stride(0), c["attn_lse"].stride(2), c["attn_lse"].stride(1),\n            c["o"].stride(0), c["o"].stride(1),\n            MAYBE_FINAL_OUT=True,\n            BATCH_NUM=batch_size,\n            BLOCK_DV=BLOCK_DV,\n            Lv=Lv,\n            mgc=64,\n            num_warps=4,\n            num_stages=2,\n            waves_per_eu=4,\n        )\n        return c["o"]\n\n    # ---- Tier 2: ALL fp8 shapes -> persistent (s3-s8) ----\n    # Non-persistent fp8 was faster but fails leaderboard correctness (v77, v78).\n    # Persistent + mla_reduce_v1 is the only leaderboard-safe fp8 path.\n    else:\n        kv_buffer_fp8, kv_scale = kv_data["fp8"]\n        kv_buffer_4d = kv_buffer_fp8.view(total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)\n\n        # Per-shape split tuning\n        if total_kv >= 1000000:\n            # s8 (256, 8192) -- splits=4 (from v77)\n            splits, fast_mode = 4, False\n        elif total_kv >= 300000:\n            # s6 (64, 8192) -- splits=4 (from 8, 64*4=256 programs = 100% CU fill)\n            splits, fast_mode = 4, False\n        elif batch_size >= 256:\n            # s7 (256, 1024) -- splits=4 with kv_gran=64 (v131 LB-safe config)\n            # splits=1+kv_gran=64 FAILED LB in v136. splits=4 gives 1024 programs.\n            splits, fast_mode = 4, False\n        elif batch_size >= 64:\n            # s5 (64, 1024) -- splits=2 (from 4, reduce=7us → ~3.5us, 128 programs = 50% CU)\n            splits, fast_mode = 2, False\n        else:\n            # s3 (32, 1024) and s4 (32, 8192)\n            if kv_seq_len <= 1024:\n                splits, fast_mode = 4, True   # s3: reduced from 8 to 4\n            else:\n                splits, fast_mode = 32, True  # s4\n\n        # Use kv_granularity=64 for ALL persistent shapes (matches v131 LB-safe config)\n        # v131 (kv_gran=64 all) PASSED LB at 58.0us. v137/v138 (kv_gran=16 for s3/s5)\n        # FAILED LB on s3. kv_gran=64 is required for LB correctness.\n        kv_gran = 64\n        c = _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, splits, fast_mode, kv_gran)\n\n        # Fast FP8 quant: copy_ cast (scale=1.0) -- from v63\n        q_2d = q.view(total_q, NUM_HEADS * QK_HEAD_DIM)\n        c["q_fp8"].copy_(q_2d)\n\n        aiter.mla_decode_stage1_asm_fwd(\n            c["q_fp8"].view(-1, NUM_HEADS, QK_HEAD_DIM), kv_buffer_4d,\n            qo_indptr, kv_indptr, c["kv_indices"], c["kv_last_page_len"],\n            None, c["work_metadata"], c["work_indptr"], c["work_info_set"],\n            1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,\n            c["logits"], c["attn_lse"], c["o"],\n            c["q_scale"], kv_scale,\n        )\n\n        aiter.mla_reduce_v1(\n            c["logits"], c["attn_lse"],\n            c["reduce_indptr"], c["reduce_final_map"], c["reduce_partial_map"],\n            1, c["o"], None,\n        )\n        return c["o"]\n'

def _load_module(name, source):
    tmpdir = tempfile.mkdtemp()
    path = os.path.join(tmpdir, name + '.py')
    open(path, 'w').write(source)
    spec = importlib.util.spec_from_file_location(name, path)
    mod = importlib.util.module_from_spec(spec)
    sys.modules[name] = mod
    spec.loader.exec_module(mod)
    return mod

# LAZY loading: modules are loaded on first use, not at import time.
# This prevents aiter (reference) from polluting Triton (target) state.
_target_mod = None
_ref_mod = None

from task import input_t, output_t

def custom_kernel(data: input_t) -> output_t:
    global _target_mod, _ref_mod
    _q, _cfg = data[0], data[4]
    if (_q.shape[0] == 32 and _cfg["kv_seq_len"] == 8192):
        if _target_mod is None:
            _target_mod = _load_module('_bss_target', _TARGET_SOURCE)
        return _target_mod.custom_kernel(data)
    else:
        if _ref_mod is None:
            _ref_mod = _load_module('_bss_ref', _REF_SOURCE)
        return _ref_mod.custom_kernel(data)
scrolls · 41 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