Skip to content
KernelIndex
Search⌘K

submission 660305

yuzhou_lithos · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v96_persist_split1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-660305?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
#223 of 766
2026-03-29

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:68481b60e0fad5b80a16725003a84a9c50b4eab14382fb6d0367e6cd3aff427a
license declaredunknown
license concludedunknown
authorsyuzhou_lithos
imported2026-08-15

Techniques

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

persistent-kernel"""MLA decode v96 — persistent with max_split_per_batch=1 for bs=256 — Wave 248.

Kernel source

submission_v96_persist_split1.py140 lines
"""MLA decode v96 — persistent with max_split_per_batch=1 for bs=256 — Wave 248.

v94's non-persistent 1-split was unreliable (correctness failures in leaderboard).
This version stays FULLY persistent but uses max_split_per_batch=1 for bs=256,
minimizing the reduce work while keeping the proven .co kernel.

For bs=256: persistent with max_split=1 → 256 CUs, 1 batch/CU, minimal reduce.
For other bs: persistent with max_split=32 (v82 approach, proven correct).
"""
import math, torch
from task import input_t, output_t

import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

SM_SCALE = 1.0 / math.sqrt(576)
FP8_DTYPE = aiter_dtypes.fp8
PAGE_SIZE = 1
H = 16
NKV = 1
DK = 576
DV = 512

_cache = {}


def _get_splits(bs):
    """Choose max_split_per_batch based on batch size."""
    if bs >= 256:
        return 1   # 256 batches fully fill 256 CUs, no KV splitting needed
    return 32      # default: 32 max splits for CU utilization


def _ensure_cached(bs, kvl, device):
    key = (bs, kvl)
    if key in _cache:
        return _cache[key]

    total_kv = bs * kvl
    max_splits = _get_splits(bs)

    qo_indptr = torch.arange(0, bs + 1, dtype=torch.int32, device=device)
    kv_indptr = torch.arange(0, bs + 1, dtype=torch.int32, device=device) * kvl
    kv_last_page_len = torch.full((bs,), kvl, dtype=torch.int32, device=device)
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)

    out = torch.empty((bs, H, DV), dtype=torch.bfloat16, device=device)
    q_fp8 = torch.empty((bs * H, DK), dtype=FP8_DTYPE, device=device)
    q_fp8_3d = q_fp8.view(bs, H, DK)
    q_scale = torch.ones(1, dtype=torch.float32, device=device)

    info = get_mla_metadata_info_v1(
        bs, 1, H, FP8_DTYPE, FP8_DTYPE,
        is_sparse=False, fast_mode=False,
        num_kv_splits=max_splits, intra_batch_mode=True,
    )
    work = [torch.empty(s, dtype=t, device=device) 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,
        H // 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, 16),
        max_seqlen_qo=1,
        uni_seqlen_qo=1,
        fast_mode=False,
        max_split_per_batch=max_splits,
        intra_batch_mode=True,
        dtype_q=FP8_DTYPE,
        dtype_kv=FP8_DTYPE,
    )

    rp_size = reduce_partial_map.size(0)
    logits = torch.empty((rp_size, 1, H, DV), dtype=torch.float32, device=device)
    attn_lse = torch.empty((rp_size, 1, H, 1), dtype=torch.float32, device=device)

    _cache[key] = {
        "qo_indptr": qo_indptr,
        "kv_indptr": kv_indptr,
        "kv_indices": kv_indices,
        "kv_last_page_len": kv_last_page_len,
        "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,
        "output": out,
        "q_fp8": q_fp8,
        "q_fp8_3d": q_fp8_3d,
        "q_scale": q_scale,
        "logits": logits,
        "attn_lse": attn_lse,
    }
    return _cache[key]


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr_in, kv_indptr_in, cfg = data
    bs = cfg["batch_size"]
    kvl = cfg["kv_seq_len"]

    c = _ensure_cached(bs, kvl, q.device)

    q_fp8 = c["q_fp8"]
    q_fp8_3d = c["q_fp8_3d"]
    logits = c["logits"]
    attn_lse = c["attn_lse"]
    output = c["output"]
    q_scale = c["q_scale"]

    q_2d = q.view(-1, DK)
    q_fp8.copy_(q_2d)

    kv_fp8, kv_scale = kv_data["fp8"]
    tkv = bs * kvl
    kv_4d = kv_fp8[:tkv].view(tkv, PAGE_SIZE, NKV, DK)

    aiter.mla_decode_stage1_asm_fwd(
        q_fp8_3d, kv_4d,
        c["qo_indptr"], c["kv_indptr"], c["kv_indices"], c["kv_last_page_len"],
        None, c["work_meta_data"], c["work_indptr"], c["work_info_set"],
        1, PAGE_SIZE, NKV, SM_SCALE,
        logits, attn_lse, output, q_scale, kv_scale,
    )

    aiter.mla_reduce_v1(
        logits, attn_lse,
        c["reduce_indptr"], c["reduce_final_map"], c["reduce_partial_map"],
        1, output, None,
    )

    return output
scrolls · 140 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