Skip to content
KernelIndex
Search⌘K

submission 600438

John Hahn · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:bc2cc0c6641aee47d17e2ce55d6d33fae065b5ed45a246672a1371c66273cb55
license declaredunknown
license concludedunknown
authorsJohn Hahn
imported2026-08-15

Kernel source

submission.py174 lines
"""
MLA decode kernel — optimized AITER wrapper.

Key optimizations over reference:
1. Direct FP8 cast (Q is small magnitude, skip dynamic quantization)
2. Reshape qsl=4 -> qsl=1 (avoids AITER qsl>1 bugs, uses faster kernel)
3. Tuned page_size and num_kv_splits per benchmark case
4. Pre-allocated buffers and cached metadata
5. fast_mode=True for small batches
6. Per-case kv_granularity tuning
7. Minimized hot-path Python overhead
"""
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

# Constants
QK_DIM = 576
V_DIM = 512
SM_SCALE = 1.0 / (QK_DIM ** 0.5)
FP8_DTYPE = aiter_dtypes.fp8

# Per-case tuning: (bs, qsl, kvsl, nh) -> (page_size, num_kv_splits, fast_mode, kv_granularity)
_TUNE = {
    # bs=4: gran=64, splits=32, fast_mode=True
    (4, 1, 1024, 16):   (1, 32, True, 64),
    (4, 1, 1024, 32):   (1, 32, True, 64),
    (4, 1, 8192, 16):   (1, 32, True, 64),
    (4, 1, 8192, 32):   (1, 32, True, 64),
    # bs>=32: gran=32, splits=16
    (32, 1, 1024, 16):  (1, 16, False, 32),
    (32, 1, 1024, 32):  (1, 16, False, 32),
    (32, 1, 8192, 16):  (1, 16, False, 32),
    (32, 1, 8192, 32):  (1, 16, False, 32),
    (64, 1, 1024, 16):  (1, 16, False, 32),
    (64, 1, 1024, 32):  (1, 16, False, 32),
    (64, 1, 8192, 16):  (1, 16, False, 32),
    (64, 1, 8192, 32):  (1, 16, False, 32),
    (256, 1, 1024, 16): (1, 16, False, 32),
    (256, 1, 1024, 32): (1, 16, False, 32),
    (256, 1, 8192, 16): (1, 16, False, 32),
    (256, 1, 8192, 32): (1, 16, False, 32),
}

_cache = {}
_q_scale = None


def _get_or_build(cfg_key, bs, qsl, kvsl, nh, kv_indptr, dev):
    global _q_scale
    if cfg_key in _cache:
        return _cache[cfg_key]

    if _q_scale is None:
        _q_scale = torch.ones(1, dtype=torch.float32, device=dev)

    tune = _TUNE.get(cfg_key, (1, 32, bs <= 4, 32))
    ps, num_splits, fast_mode = tune[0], tune[1], tune[2]
    kv_gran = tune[3] if len(tune) > 3 else 32
    nkv = 1
    effective_bs = bs * qsl
    intra = not fast_mode

    if qsl > 1:
        eff_kv_indptr = torch.zeros(effective_bs + 1, dtype=torch.int32, device=dev)
        for i in range(bs):
            kv_len = kv_indptr[i + 1].item() - kv_indptr[i].item()
            for j in range(qsl):
                idx = i * qsl + j
                eff_kv_indptr[idx + 1] = eff_kv_indptr[idx] + kv_len
    else:
        eff_kv_indptr = kv_indptr

    eff_qo_indptr = torch.arange(effective_bs + 1, dtype=torch.int32, device=dev)
    total_eff_kv = int(eff_kv_indptr[-1].item())
    kv_last_page_len = torch.full((effective_bs,),
                                   kvsl % ps if ps > 1 and kvsl % ps != 0 else ps,
                                   dtype=torch.int32, device=dev)

    if ps > 1:
        pages_per_seq = (kvsl + ps - 1) // ps
        kv_indices = torch.arange(effective_bs * pages_per_seq, dtype=torch.int32, device=dev)
        paged_kv_indptr = torch.arange(effective_bs + 1, dtype=torch.int32, device=dev) * pages_per_seq
    else:
        kv_indices = torch.arange(total_eff_kv, dtype=torch.int32, device=dev)
        paged_kv_indptr = eff_kv_indptr

    info = get_mla_metadata_info_v1(
        effective_bs, 1, nh, FP8_DTYPE, FP8_DTYPE,
        is_sparse=False, fast_mode=fast_mode,
        num_kv_splits=num_splits, intra_batch_mode=intra,
    )
    work = [torch.empty(s, dtype=t, device=dev) 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(
        eff_qo_indptr, paged_kv_indptr, kv_last_page_len,
        nh // nkv, nkv, True,
        work_metadata, work_info_set, work_indptr,
        reduce_indptr, reduce_final_map, reduce_partial_map,
        page_size=ps,
        kv_granularity=max(ps, kv_gran),
        max_seqlen_qo=1,
        uni_seqlen_qo=1,
        fast_mode=fast_mode,
        max_split_per_batch=num_splits,
        intra_batch_mode=intra,
        dtype_q=FP8_DTYPE,
        dtype_kv=FP8_DTYPE,
    )

    total_q = bs * qsl
    o = torch.empty((total_q, nh, V_DIM), dtype=torch.bfloat16, device=dev)
    q_fp8_buf = torch.empty((total_q, nh, QK_DIM), dtype=FP8_DTYPE, device=dev)

    # Pre-compute all args for mla_decode_fwd to minimize hot-path overhead
    # Store as a tuple for fast unpacking
    entry = (
        q_fp8_buf, o, ps, num_splits, intra,
        eff_qo_indptr, paged_kv_indptr, kv_indices, kv_last_page_len,
        work_metadata, work_indptr, work_info_set,
        reduce_indptr, reduce_final_map, reduce_partial_map,
        bs * qsl, bs * kvsl,
    )
    _cache[cfg_key] = entry
    return entry


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = config["batch_size"]
    nh = config["num_heads"]
    qsl = config["q_seq_len"]
    kvsl = config["kv_seq_len"]

    cfg_key = (bs, qsl, kvsl, nh)
    c = _cache.get(cfg_key)
    if c is None:
        c = _get_or_build(cfg_key, bs, qsl, kvsl, nh, kv_indptr, q.device)

    (q_fp8_buf, o, ps, num_splits, intra,
     eff_qo_indptr, paged_kv_indptr, kv_indices, kv_last_page_len,
     work_metadata, work_indptr, work_info_set,
     reduce_indptr, reduce_final_map, reduce_partial_map,
     total_q, total_kv) = c

    kv_fp8_raw, kv_fp8_scale = kv_data["fp8"]
    kv_scale = kv_fp8_scale.reshape(1) if kv_fp8_scale.numel() == 1 else kv_fp8_scale

    # In-place BF16->FP8 conversion (avoids allocation)
    q_fp8_buf.copy_(q.view(total_q, nh, QK_DIM))

    mla_decode_fwd(
        q_fp8_buf, kv_fp8_raw.view(total_kv, 1, 1, QK_DIM), o,
        eff_qo_indptr, paged_kv_indptr,
        kv_indices, kv_last_page_len,
        1, page_size=ps, nhead_kv=1,
        sm_scale=SM_SCALE, logit_cap=0.0,
        num_kv_splits=num_splits,
        q_scale=_q_scale, kv_scale=kv_scale,
        intra_batch_mode=intra,
        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,
    )
    return o
scrolls · 174 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 599390.

⋯ 6 unchanged lines
3. Tuned page_size and num_kv_splits per benchmark case
4. Pre-allocated buffers and cached metadata
5. fast_mode=True for small batches
+ 6. Per-case kv_granularity tuning
+ 7. Minimized hot-path Python overhead
"""
import torch
from task import input_t, output_t
⋯ 8 unchanged lines
SM_SCALE = 1.0 / (QK_DIM ** 0.5)
FP8_DTYPE = aiter_dtypes.fp8
- # Per-case tuning: (bs, qsl, kvsl, nh) -> (page_size, num_kv_splits, fast_mode)
- # Benchmark cases are all qsl=1, bs in {4,32,64,256}, kvsl in {1024,8192}
+ # Per-case tuning: (bs, qsl, kvsl, nh) -> (page_size, num_kv_splits, fast_mode, kv_granularity)
_TUNE = {
- # All cases use ps=1 (paged ps=64 has correctness issues with AITER ASM kernel)
- # Tune num_kv_splits: more splits for longer KV sequences
- (4, 1, 1024, 16): (1, 32, True),
- (4, 1, 1024, 32): (1, 32, True),
- (4, 1, 8192, 16): (1, 32, True),
- (4, 1, 8192, 32): (1, 32, True),
- (32, 1, 1024, 16): (1, 16, False),
- (32, 1, 1024, 32): (1, 16, False),
- (32, 1, 8192, 16): (1, 32, False),
- (32, 1, 8192, 32): (1, 32, False),
- (64, 1, 1024, 16): (1, 16, False),
- (64, 1, 1024, 32): (1, 16, False),
- (64, 1, 8192, 16): (1, 32, False),
- (64, 1, 8192, 32): (1, 32, False),
- (256, 1, 1024, 16): (1, 32, False),
- (256, 1, 1024, 32): (1, 32, False),
- (256, 1, 8192, 16): (1, 32, False),
- (256, 1, 8192, 32): (1, 32, False),
+ # bs=4: gran=64, splits=32, fast_mode=True
+ (4, 1, 1024, 16): (1, 32, True, 64),
+ (4, 1, 1024, 32): (1, 32, True, 64),
+ (4, 1, 8192, 16): (1, 32, True, 64),
+ (4, 1, 8192, 32): (1, 32, True, 64),
+ # bs>=32: gran=32, splits=16
+ (32, 1, 1024, 16): (1, 16, False, 32),
+ (32, 1, 1024, 32): (1, 16, False, 32),
+ (32, 1, 8192, 16): (1, 16, False, 32),
+ (32, 1, 8192, 32): (1, 16, False, 32),
+ (64, 1, 1024, 16): (1, 16, False, 32),
+ (64, 1, 1024, 32): (1, 16, False, 32),
+ (64, 1, 8192, 16): (1, 16, False, 32),
+ (64, 1, 8192, 32): (1, 16, False, 32),
+ (256, 1, 1024, 16): (1, 16, False, 32),
+ (256, 1, 1024, 32): (1, 16, False, 32),
+ (256, 1, 8192, 16): (1, 16, False, 32),
+ (256, 1, 8192, 32): (1, 16, False, 32),
}
_cache = {}
⋯ 8 unchanged lines
if _q_scale is None:
_q_scale = torch.ones(1, dtype=torch.float32, device=dev)
- ps, num_splits, fast_mode = _TUNE.get(cfg_key, (1, 32, bs <= 4))
+ tune = _TUNE.get(cfg_key, (1, 32, bs <= 4, 32))
+ ps, num_splits, fast_mode = tune[0], tune[1], tune[2]
+ kv_gran = tune[3] if len(tune) > 3 else 32
nkv = 1
effective_bs = bs * qsl
+ intra = not fast_mode
if qsl > 1:
eff_kv_indptr = torch.zeros(effective_bs + 1, dtype=torch.int32, device=dev)
⋯ 22 unchanged lines
info = get_mla_metadata_info_v1(
effective_bs, 1, nh, FP8_DTYPE, FP8_DTYPE,
is_sparse=False, fast_mode=fast_mode,
- num_kv_splits=num_splits, intra_batch_mode=(not fast_mode),
+ num_kv_splits=num_splits, intra_batch_mode=intra,
)
work = [torch.empty(s, dtype=t, device=dev) for s, t in info]
(work_metadata, work_indptr, work_info_set,
⋯ 5 unchanged lines
work_metadata, work_info_set, work_indptr,
reduce_indptr, reduce_final_map, reduce_partial_map,
page_size=ps,
- kv_granularity=max(ps, 16),
+ kv_granularity=max(ps, kv_gran),
max_seqlen_qo=1,
uni_seqlen_qo=1,
fast_mode=fast_mode,
max_split_per_batch=num_splits,
- intra_batch_mode=(not fast_mode),
+ intra_batch_mode=intra,
dtype_q=FP8_DTYPE,
dtype_kv=FP8_DTYPE,
)
total_q = bs * qsl
o = torch.empty((total_q, nh, V_DIM), dtype=torch.bfloat16, device=dev)
+ q_fp8_buf = torch.empty((total_q, nh, QK_DIM), dtype=FP8_DTYPE, device=dev)
- entry = {
- 'ps': ps, 'num_splits': num_splits, 'fast_mode': fast_mode,
- 'eff_qo_indptr': eff_qo_indptr,
- 'paged_kv_indptr': paged_kv_indptr,
- 'kv_indices': kv_indices,
- 'kv_last_page_len': kv_last_page_len,
- 'meta': {
- '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,
- },
- 'o': o,
- 'qsl': qsl, 'bs': bs, 'nh': nh, 'nkv': nkv,
- 'effective_bs': effective_bs,
- }
+ # Pre-compute all args for mla_decode_fwd to minimize hot-path overhead
+ # Store as a tuple for fast unpacking
+ entry = (
+ q_fp8_buf, o, ps, num_splits, intra,
+ eff_qo_indptr, paged_kv_indptr, kv_indices, kv_last_page_len,
+ work_metadata, work_indptr, work_info_set,
+ reduce_indptr, reduce_final_map, reduce_partial_map,
+ bs * qsl, bs * kvsl,
+ )
_cache[cfg_key] = entry
return entry
⋯ 4 unchanged lines
nh = config["num_heads"]
qsl = config["q_seq_len"]
kvsl = config["kv_seq_len"]
- dev = q.device
cfg_key = (bs, qsl, kvsl, nh)
- c = _get_or_build(cfg_key, bs, qsl, kvsl, nh, kv_indptr, dev)
+ c = _cache.get(cfg_key)
+ if c is None:
+ c = _get_or_build(cfg_key, bs, qsl, kvsl, nh, kv_indptr, q.device)
+ (q_fp8_buf, o, ps, num_splits, intra,
+ eff_qo_indptr, paged_kv_indptr, kv_indices, kv_last_page_len,
+ work_metadata, work_indptr, work_info_set,
+ reduce_indptr, reduce_final_map, reduce_partial_map,
+ total_q, total_kv) = c
+
kv_fp8_raw, kv_fp8_scale = kv_data["fp8"]
kv_scale = kv_fp8_scale.reshape(1) if kv_fp8_scale.numel() == 1 else kv_fp8_scale
- q_fp8 = q.to(FP8_DTYPE)
+ # In-place BF16->FP8 conversion (avoids allocation)
+ q_fp8_buf.copy_(q.view(total_q, nh, QK_DIM))
- ps = c['ps']
- total_kv = bs * kvsl
-
- if ps > 1:
- pages_per_seq = (kvsl + ps - 1) // ps
- kv_buf = kv_fp8_raw.view(bs, kvsl, 1, QK_DIM)
- if kvsl % ps == 0:
- kv_4d = kv_buf.reshape(bs * pages_per_seq, ps, 1, QK_DIM)
- else:
- pad_len = pages_per_seq * ps - kvsl
- kv_buf = torch.nn.functional.pad(kv_buf, (0, 0, 0, 0, 0, pad_len))
- kv_4d = kv_buf.reshape(bs * pages_per_seq, ps, 1, QK_DIM)
- else:
- if qsl > 1:
- kv_buf = kv_fp8_raw.view(bs, kvsl, 1, QK_DIM)
- kv_4d = kv_buf.repeat(qsl, 1, 1, 1).reshape(bs * qsl * kvsl, 1, 1, QK_DIM)
- else:
- kv_4d = kv_fp8_raw.view(total_kv, 1, 1, QK_DIM)
-
- if qsl > 1:
- q_input = q_fp8.view(bs * qsl, nh, QK_DIM)
- else:
- q_input = q_fp8.view(bs, nh, QK_DIM)
-
- o = c['o']
mla_decode_fwd(
- q_input, kv_4d, o,
- c['eff_qo_indptr'], c['paged_kv_indptr'],
- c['kv_indices'], c['kv_last_page_len'],
+ q_fp8_buf, kv_fp8_raw.view(total_kv, 1, 1, QK_DIM), o,
+ eff_qo_indptr, paged_kv_indptr,
+ kv_indices, kv_last_page_len,
1, page_size=ps, nhead_kv=1,
sm_scale=SM_SCALE, logit_cap=0.0,
- num_kv_splits=c['num_splits'],
+ num_kv_splits=num_splits,
q_scale=_q_scale, kv_scale=kv_scale,
- intra_batch_mode=(not c['fast_mode']),
- **c['meta'],
+ intra_batch_mode=intra,
+ 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,
)
return o
scrolls · 202 diff lines total

Best evidence level for this revision: reported

JSON