Skip to content
KernelIndex
Search⌘K

submission 715009

Lemonade · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:bcae29a8a241d0c8a03ac445d400ccebd274c3c9361aac63d0137a13321d923b
license declaredunknown
license concludedunknown
authorsLemonade
imported2026-08-26

Kernel source

submission.py120 lines
import torch
import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.ops.triton.quant.quant import dynamic_per_tensor_quant_fp8_i8

FP8_DTYPE = aiter_dtypes.fp8
FP8_MAX = torch.finfo(FP8_DTYPE).max
FP8_MIN = torch.finfo(FP8_DTYPE).min
SM_SCALE = 1.0 / (576.0 ** 0.5)

_cache = {}


def _ensure_cache(bs, kv_len, device, qo_indptr, kv_indptr):
    key = (bs, kv_len)
    if key in _cache:
        return _cache[key]
    total_kv = bs * kv_len; total_q = bs

    # All a8w8 — single ASM kernel type, safe for leaderboard
    q_fp8 = torch.empty((total_q * 16, 576), dtype=FP8_DTYPE, device=device)
    q_scale = torch.empty(1, dtype=torch.float32, device=device)

    if bs >= 128:
        # Non-persistent 1-split: skip reduce entirely (no Triton stage2 needed)
        entry = {
            'mode': 'np1', 'q_fp8': q_fp8, 'q_scale': q_scale,
            'nksi': torch.arange(0, bs + 1, dtype=torch.int32, device=device),
            'kv_last': torch.full((bs,), kv_len, dtype=torch.int32, device=device),
            'kv_idx': torch.arange(total_kv, dtype=torch.int32, device=device),
            'lse': torch.empty((total_q, 1, 16, 1), dtype=torch.float32, device=device),
        }
    elif bs >= 64:
        # Persistent for bs=64 (better CU utilization than np1)
        ns = min(max(2, min(256, 512 // bs)), max(1, kv_len // 16))
        kv_last = torch.full((bs,), kv_len, dtype=torch.int32, device=device)
        kv_idx = torch.arange(total_kv, dtype=torch.int32, device=device)
        info = get_mla_metadata_info_v1(bs, 1, 16, FP8_DTYPE, FP8_DTYPE,
            is_sparse=False, fast_mode=True, num_kv_splits=ns, intra_batch_mode=True)
        work = [torch.empty(s, dtype=t, device=device) for s, t in info]
        wm, wi, wis, ri, rfm, rpm = work
        get_mla_metadata_v1(qo_indptr, kv_indptr, kv_last, 16, 1, True,
            wm, wis, wi, ri, rfm, rpm, page_size=1, kv_granularity=16,
            max_seqlen_qo=1, uni_seqlen_qo=1, fast_mode=True,
            max_split_per_batch=ns, intra_batch_mode=True,
            dtype_q=FP8_DTYPE, dtype_kv=FP8_DTYPE)
        rpm_size = rpm.shape[0]
        entry = {
            'mode': 'ps', 'q_fp8': q_fp8, 'q_scale': q_scale,
            'kv_last': kv_last, 'kv_idx': kv_idx,
            'wm': wm, 'wi': wi, 'wis': wis, 'ri': ri, 'rfm': rfm, 'rpm': rpm,
            'logits': torch.empty((rpm_size, 1, 16, 512), dtype=torch.float32, device=device),
            'lse': torch.empty((rpm_size, 1, 16, 1), dtype=torch.float32, device=device),
        }
    else:
        # Persistent with CU-filling splits
        ns = min(max(2, min(256, 512 // bs)), max(1, kv_len // 16))
        kv_last = torch.full((bs,), kv_len, dtype=torch.int32, device=device)
        kv_idx = torch.arange(total_kv, dtype=torch.int32, device=device)
        info = get_mla_metadata_info_v1(bs, 1, 16, FP8_DTYPE, FP8_DTYPE,
            is_sparse=False, fast_mode=True, num_kv_splits=ns, intra_batch_mode=True)
        work = [torch.empty(s, dtype=t, device=device) for s, t in info]
        wm, wi, wis, ri, rfm, rpm = work
        get_mla_metadata_v1(qo_indptr, kv_indptr, kv_last, 16, 1, True,
            wm, wis, wi, ri, rfm, rpm, page_size=1, kv_granularity=16,
            max_seqlen_qo=1, uni_seqlen_qo=1, fast_mode=True,
            max_split_per_batch=ns, intra_batch_mode=True,
            dtype_q=FP8_DTYPE, dtype_kv=FP8_DTYPE)
        rpm_size = rpm.shape[0]
        entry = {
            'mode': 'ps', 'q_fp8': q_fp8, 'q_scale': q_scale,
            'kv_last': kv_last, 'kv_idx': kv_idx,
            'wm': wm, 'wi': wi, 'wis': wis, 'ri': ri, 'rfm': rfm, 'rpm': rpm,
            'logits': torch.empty((rpm_size, 1, 16, 512), dtype=torch.float32, device=device),
            'lse': torch.empty((rpm_size, 1, 16, 1), dtype=torch.float32, device=device),
        }
    _cache[key] = entry
    return entry


def custom_kernel(data):
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = config["batch_size"]
    kv_len = config["kv_seq_len"]

    c = _ensure_cache(bs, kv_len, q.device, qo_indptr, kv_indptr)

    # fp8 Q quantization — hybrid: fused Triton for small, static scale for large
    q_2d = q.view(-1, 576)
    if bs <= 32:
        dynamic_per_tensor_quant_fp8_i8(c['q_fp8'], q_2d, c['q_scale'])
    else:
        FIXED_SCALE = 6.0 / FP8_MAX
        c['q_fp8'].copy_((q_2d * (1.0 / FIXED_SCALE)).clamp(min=FP8_MIN, max=FP8_MAX).to(FP8_DTYPE))
        c['q_scale'].fill_(FIXED_SCALE)

    q3 = c['q_fp8'].view(bs, 16, 576)
    kv_fp8, kv_scale = kv_data["fp8"]
    kv_4d = kv_fp8.view(-1, 1, 1, 576)
    o = torch.empty(q.shape[0], 16, 512, dtype=torch.bfloat16, device=q.device)
    mode = c['mode']

    if mode == 'np1':
        lv = o.view(bs, 1, 16, 512)
        aiter.mla_decode_stage1_asm_fwd(
            q3, kv_4d, qo_indptr, kv_indptr, c['kv_idx'], c['kv_last'],
            c['nksi'], None, None, None, 1, 1, 1, SM_SCALE,
            lv, c['lse'], o, c['q_scale'], kv_scale)
    else:
        # Persistent mode (bs<128)
        aiter.mla_decode_stage1_asm_fwd(
            q3, kv_4d, qo_indptr, kv_indptr, c['kv_idx'], c['kv_last'],
            None, c['wm'], c['wi'], c['wis'], 1, 1, 1, SM_SCALE,
            c['logits'], c['lse'], o, c['q_scale'], kv_scale)
        aiter.mla_reduce_v1(c['logits'], c['lse'],
            c['ri'], c['rfm'], c['rpm'], 1, o, None)

    return o
scrolls · 120 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