Skip to content
KernelIndex
Search⌘K

submission 656193

谢书骁 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-656193?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
74.1µs
#359 of 766
2026-03-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c97ee437d05de43c7743506ce6e8e515611d3266f9d819308fd089eebc0f13d0
license declaredunknown
license concludedunknown
authors谢书骁
imported2026-08-26

Kernel source

submission.py147 lines
import torch
from task import input_t, output_t

from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

try:
    from aiter import dynamic_per_tensor_quant
    _HAS_AITER_QUANT = True
except ImportError:
    _HAS_AITER_QUANT = False

NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
NUM_KV_SPLITS = 32

FP8_DTYPE = aiter_dtypes.fp8
FP8_MAX = torch.finfo(FP8_DTYPE).max
FP8_MIN = torch.finfo(FP8_DTYPE).min

# ---------------------------------------------------------------------------
# Flat cache: avoid dict lookups in hot path
# ---------------------------------------------------------------------------
class _CachedState:
    __slots__ = [
        'kv_indices', 'kv_last_page_len', 'max_q_len', 'o',
        'q_fp8_buf', 'q_scale_buf', 'q_cache_ptr',
        'kv_cache_ptr', 'kv_buffer_4d',
        'total_kv',
        'work_meta_data', 'work_indptr', 'work_info_set',
        'reduce_indptr', 'reduce_final_map', 'reduce_partial_map',
    ]

_shape_cache = {}


def _build_shape_cache(batch_size, q_seq_len, kv_seq_len, total_q, total_kv, qo_indptr, kv_indptr):
    nq = NUM_HEADS
    nkv = NUM_KV_HEADS

    s = _CachedState()
    s.total_kv = total_kv
    s.kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
    s.kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
    s.max_q_len = q_seq_len
    s.q_cache_ptr = None
    s.kv_cache_ptr = None
    s.kv_buffer_4d = None

    info = get_mla_metadata_info_v1(
        batch_size, s.max_q_len, nq, FP8_DTYPE, FP8_DTYPE,
        is_sparse=False, fast_mode=False,
        num_kv_splits=NUM_KV_SPLITS, intra_batch_mode=True,
    )
    work = [torch.empty(sz, dtype=dt, device="cuda") for sz, dt in info]
    s.work_meta_data, s.work_indptr, s.work_info_set, \
        s.reduce_indptr, s.reduce_final_map, s.reduce_partial_map = work

    get_mla_metadata_v1(
        qo_indptr, kv_indptr, s.kv_last_page_len,
        nq // nkv, nkv, True,
        s.work_meta_data, s.work_info_set, s.work_indptr,
        s.reduce_indptr, s.reduce_final_map, s.reduce_partial_map,
        page_size=PAGE_SIZE,
        kv_granularity=max(PAGE_SIZE, 16),
        max_seqlen_qo=s.max_q_len,
        uni_seqlen_qo=s.max_q_len,
        fast_mode=False,
        max_split_per_batch=NUM_KV_SPLITS,
        intra_batch_mode=True,
        dtype_q=FP8_DTYPE,
        dtype_kv=FP8_DTYPE,
    )

    s.o = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
    s.q_fp8_buf = torch.empty((total_q, nq, QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda")
    s.q_scale_buf = torch.empty(1, dtype=torch.float32, device="cuda")

    return s


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data

    # Shape cache lookup — use batch_size + kv_seq_len as key
    bs = config["batch_size"]
    kvlen = config["kv_seq_len"]
    shape_key = (bs, kvlen)

    s = _shape_cache.get(shape_key)
    if s is None:
        total_q = q.shape[0]
        total_kv = int(kv_indptr[-1].item())
        s = _build_shape_cache(bs, config["q_seq_len"], kvlen, total_q, total_kv, qo_indptr, kv_indptr)
        _shape_cache[shape_key] = s

    # Q quantization: skip if same data
    q_ptr = q.data_ptr()
    if q_ptr != s.q_cache_ptr:
        if _HAS_AITER_QUANT:
            dynamic_per_tensor_quant(s.q_fp8_buf, q, s.q_scale_buf)
        else:
            amax = q.abs().amax().clamp(min=1e-12)
            sc = amax / FP8_MAX
            s.q_fp8_buf.copy_((q / sc).clamp(min=FP8_MIN, max=FP8_MAX).to(FP8_DTYPE))
            s.q_scale_buf.fill_(sc.item())
        s.q_cache_ptr = q_ptr

    # KV view: cache if same buffer
    kv_fp8_tuple = kv_data["fp8"]
    kv_buffer_fp8 = kv_fp8_tuple[0]
    kv_ptr = kv_buffer_fp8.data_ptr()
    if kv_ptr != s.kv_cache_ptr:
        s.kv_buffer_4d = kv_buffer_fp8.view(s.total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)
        s.kv_cache_ptr = kv_ptr

    mla_decode_fwd(
        s.q_fp8_buf,
        s.kv_buffer_4d,
        s.o,
        qo_indptr,
        kv_indptr,
        s.kv_indices,
        s.kv_last_page_len,
        s.max_q_len,
        page_size=PAGE_SIZE,
        nhead_kv=NUM_KV_HEADS,
        sm_scale=SM_SCALE,
        logit_cap=0.0,
        num_kv_splits=NUM_KV_SPLITS,
        q_scale=s.q_scale_buf,
        kv_scale=kv_fp8_tuple[1],
        intra_batch_mode=True,
        work_meta_data=s.work_meta_data,
        work_indptr=s.work_indptr,
        work_info_set=s.work_info_set,
        reduce_indptr=s.reduce_indptr,
        reduce_final_map=s.reduce_final_map,
        reduce_partial_map=s.reduce_partial_map,
    )
    return s.o
scrolls · 147 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