Skip to content
KernelIndex
Search⌘K

submission 683270

Aniket Sadashiva · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_probe_v362_morefp8_p416_shapeamax13.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-683270?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
55.6µs
#190 of 766
2026-03-31

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d7f6d8dcaef6dc4d3b6e30be0cd5e1c6881abb03e8aeec38db9872aadf1bbaa1
license declaredunknown
license concludedunknown
authorsAniket Sadashiva
imported2026-08-15

Techniques

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

fp8tl.store(out_ptr + offsets, x.to(tl.float8e4nv), mask=mask)

Kernel source

submission_probe_v362_morefp8_p416_shapeamax13.py241 lines
from task import input_t, output_t
import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (576 ** 0.5)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
FP8_MAX = torch.finfo(FP8_DTYPE).max
FP8_MIN = torch.finfo(FP8_DTYPE).min
DEFAULT_Q_AMAX = 0.12
Q_AMAX_TABLE = {
    (32, 1024): 0.13,
}

BF16_NP_SHAPES = {(4, 1024),}

FP8_SPLITS_TABLE = {
    (4, 8192): 16,
    (32, 1024): 4,
    (32, 8192): 32,
    (64, 1024): 16,
    (64, 8192): 32,
    (256, 1024): 16,
    (256, 8192): 32,
}

_cache = {}
_stage1 = None
_reduce = None
_q_scale_cache = {}
_q_fp8_bufs = {}


def _compute_np_splits(bs, kv_seq_len):
    cu_num = 304
    overhead = 84.1
    best_score = -1
    best_i = 1
    for i in range(1, 17):
        waves = ((bs * i + cu_num - 1) // cu_num) * cu_num
        score = (bs * i / waves) * kv_seq_len / (kv_seq_len + overhead * i)
        if score > best_score:
            best_score = score
            best_i = i
    return best_i


@triton.jit
def _reduce_kernel(sd_ptr, sl_ptr, o_ptr, num_valid, NUM_SPLITS: tl.constexpr):
    bid = tl.program_id(0)
    hid = tl.program_id(1)
    offs_d = tl.arange(0, 512)
    sd_base = (bid * NUM_SPLITS * 16 * 512 + hid * 512).to(tl.int64)
    sl_base = (bid * NUM_SPLITS * 16 + hid).to(tl.int64)
    e_max = -float("inf")
    e_sum = 0.0
    acc = tl.zeros((512,), dtype=tl.float32)
    for s in range(NUM_SPLITS):
        if s < num_valid:
            v = tl.load(sd_ptr + sd_base + s * 16 * 512 + offs_d)
            lse = tl.load(sl_ptr + sl_base + s * 16)
            new_max = tl.maximum(e_max, lse)
            old_scale = tl.exp(e_max - new_max)
            new_scale = tl.exp(lse - new_max)
            acc = acc * old_scale + new_scale * v
            e_sum = e_sum * old_scale + new_scale
            e_max = new_max
    tl.store(o_ptr + (bid * 16 * 512 + hid * 512).to(tl.int64) + offs_d, (acc / e_sum).to(tl.bfloat16))


@triton.jit
def _quant_q_kernel(
    q_ptr, out_ptr, n_elements, inv_scale,
    fp8_min: tl.constexpr, fp8_max: tl.constexpr, block_size: tl.constexpr,
):
    pid = tl.program_id(0)
    offsets = pid * block_size + tl.arange(0, block_size)
    mask = offsets < n_elements
    x = tl.load(q_ptr + offsets, mask=mask).to(tl.float32)
    x = x * inv_scale
    x = tl.minimum(tl.maximum(x, fp8_min), fp8_max)
    tl.store(out_ptr + offsets, x.to(tl.float8e4nv), mask=mask)


def _get_bf16_np_cached(batch_size, kv_seq_len):
    key = ("bf16np", batch_size, kv_seq_len)
    if key in _cache:
        return _cache[key]
    total_kv = batch_size * kv_seq_len
    ns = _compute_np_splits(batch_size, kv_seq_len)
    c = {
        "output": torch.empty((batch_size, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
        "qo_indptr": torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda"),
        "kv_indptr": torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda") * kv_seq_len,
        "kv_indices": torch.arange(total_kv, dtype=torch.int32, device="cuda"),
        "kv_last_page_len": torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda"),
        "ns": ns,
        "num_kv_splits_indptr": torch.arange(0, (batch_size + 1) * ns, ns, dtype=torch.int32, device="cuda"),
        "logits": torch.empty((batch_size, ns, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device="cuda"),
        "attn_lse": torch.empty((batch_size, ns, NUM_HEADS, 1), dtype=torch.float32, device="cuda"),
    }
    _cache[key] = c
    return c


def _get_fp8_ps_cached(batch_size, kv_seq_len, num_splits):
    key = ("fp8ps", batch_size, kv_seq_len, num_splits)
    if key in _cache:
        return _cache[key]

    total_kv = batch_size * kv_seq_len
    qo_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda")
    kv_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda") * kv_seq_len
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
    kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")

    info = get_mla_metadata_info_v1(
        batch_size, 1, NUM_HEADS, FP8_DTYPE, FP8_DTYPE,
        is_sparse=False, fast_mode=False,
        num_kv_splits=num_splits, intra_batch_mode=True,
    )
    work = [torch.empty(s, dtype=t, device="cuda") 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,
        NUM_HEADS // NUM_KV_HEADS, NUM_KV_HEADS, 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=num_splits, intra_batch_mode=True,
        dtype_q=FP8_DTYPE, dtype_kv=FP8_DTYPE,
    )

    logits = torch.empty(
        (reduce_partial_map.size(0), 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device="cuda",
    )
    attn_lse = torch.empty(
        (reduce_partial_map.size(0), 1, NUM_HEADS, 1), dtype=torch.float32, device="cuda",
    )

    c = {
        "output": torch.empty((batch_size, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
        "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,
        "logits": logits,
        "attn_lse": attn_lse,
    }
    _cache[key] = c
    return c


def custom_kernel(data: input_t) -> output_t:
    global _stage1, _reduce, _q_scale_cache

    q, kv_data, _, _, config = data
    batch_size = int(config["batch_size"])
    kv_seq_len = int(config["kv_seq_len"])

    if _stage1 is None:
        _stage1 = aiter.mla_decode_stage1_asm_fwd
        _reduce = aiter.mla_reduce_v1

    if (batch_size, kv_seq_len) in BF16_NP_SHAPES:
        c = _get_bf16_np_cached(batch_size, kv_seq_len)
        kv_4d = kv_data["bf16"].view(-1, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)
        _stage1(
            q.view(-1, NUM_HEADS, QK_HEAD_DIM), kv_4d,
            c["qo_indptr"], c["kv_indptr"], c["kv_indices"], c["kv_last_page_len"],
            c["num_kv_splits_indptr"], None, None, None,
            1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
            c["logits"], c["attn_lse"], c["output"],
            q_scale=None, kv_scale=None,
        )
        ns = c["ns"]
        _reduce_kernel[(batch_size, NUM_HEADS)](
            c["logits"].view(batch_size * ns, NUM_HEADS, V_HEAD_DIM),
            c["attn_lse"].view(batch_size * ns, NUM_HEADS),
            c["output"], ns, NUM_SPLITS=ns,
        )
        return c["output"]

    if batch_size not in _q_fp8_bufs:
        _q_fp8_bufs[batch_size] = torch.empty(
            (batch_size, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda"
        )

    q_fp8 = _q_fp8_bufs[batch_size]
    n_elements = batch_size * NUM_HEADS * QK_HEAD_DIM
    q_amax = Q_AMAX_TABLE.get((batch_size, kv_seq_len), DEFAULT_Q_AMAX)
    q_scale = _q_scale_cache.get(q_amax)
    if q_scale is None:
        q_scale = torch.tensor([q_amax / FP8_MAX], dtype=torch.float32, device="cuda")
        _q_scale_cache[q_amax] = q_scale
    _quant_q_kernel[((n_elements + 1023) // 1024,)](
        q, q_fp8, n_elements, FP8_MAX / q_amax,
        fp8_min=FP8_MIN, fp8_max=FP8_MAX, block_size=1024,
    )

    kv_buffer_fp8, kv_scale = kv_data["fp8"]
    num_splits = FP8_SPLITS_TABLE[(batch_size, kv_seq_len)]
    c = _get_fp8_ps_cached(batch_size, kv_seq_len, num_splits)

    _stage1(
        q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM),
        kv_buffer_fp8.view(-1, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM),
        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, NUM_KV_HEADS, SM_SCALE,
        c["logits"], c["attn_lse"], c["output"],
        q_scale=q_scale, kv_scale=kv_scale,
    )

    _reduce(
        c["logits"], c["attn_lse"],
        c["reduce_indptr"], c["reduce_final_map"], c["reduce_partial_map"],
        1, c["output"], None,
    )

    return c["output"]
scrolls · 241 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