Skip to content
KernelIndex
Search⌘K

submission 738884

Behzod12312121 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_quant_hip_splitmap_s3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-738884?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
75.6µs
#373 of 766
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:4e193b564639ef7894c730f9beb51ea542dc799d744ef20958f60d80c383b768
license declaredunknown
license concludedunknown
authorsBehzod12312121
imported2026-08-26

Techniques

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

persistent-kernelsurface. The persistent decode path and the best measured split schedule stay

Kernel source

submission_quant_hip_splitmap_s3.py302 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""
Mixed-MLA contingency candidate: use AITER's per_tensor_quant_hip helper.

This tests whether the library's higher-level quant helper performs differently
enough from the direct compiled-op wrapper to matter on the official benchmark
surface. The persistent decode path and the best measured split schedule stay
the same as the current MLA best file.
"""

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
from aiter.ops.quant import per_tensor_quant_hip


NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM
V_HEAD_DIM = KV_LORA_RANK
SM_SCALE = 1.0 / (QK_HEAD_DIM**0.5)

PAGE_SIZE = 1
NUM_KV_SPLITS = 32

FP8_DTYPE = aiter_dtypes.fp8

SPLIT_OVERRIDES = {
    (1, 512): 8,
    (1, 2048): 12,
    (1, 8192): 20,
    (8, 512): 8,
    (8, 2048): 16,
    (8, 8192): 24,
    (16, 512): 12,
    (16, 2048): 20,
    (16, 8192): 28,
    (32, 512): 16,
    (32, 2048): 24,
    (32, 8192): 32,
    (64, 512): 20,
    (64, 2048): 28,
    (64, 8192): 36,
}

KV_GRANULARITY = 32

_WORKSPACE_CACHE: dict[tuple[object, ...], tuple[torch.Tensor, ...]] = {}
_INDPTR_CACHE: dict[tuple[object, ...], tuple[torch.Tensor, torch.Tensor, torch.Tensor]] = {}
_KV_INDEX_CACHE: dict[tuple[object, ...], torch.Tensor] = {}
_METADATA_CACHE: dict[tuple[object, ...], dict[str, torch.Tensor]] = {}
_OUTPUT_CACHE: dict[tuple[object, ...], torch.Tensor] = {}


def quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    if not tensor.is_contiguous():
        tensor = tensor.contiguous()
    return per_tensor_quant_hip(tensor, quant_dtype=FP8_DTYPE)


def _device_cache_key(device: torch.device) -> tuple[str, int]:
    return device.type, -1 if device.index is None else device.index


def _get_metadata_workspace(
    batch_size: int,
    max_q_len: int,
    nhead: int,
    nhead_kv: int,
    q_dtype: torch.dtype,
    kv_dtype: torch.dtype,
    num_kv_splits: int,
    device: torch.device,
) -> tuple[torch.Tensor, ...]:
    cache_key = (
        _device_cache_key(device),
        batch_size,
        max_q_len,
        nhead,
        nhead_kv,
        q_dtype,
        kv_dtype,
        num_kv_splits,
    )
    cached = _WORKSPACE_CACHE.get(cache_key)
    if cached is None:
        info = get_mla_metadata_info_v1(
            batch_size,
            max_q_len,
            nhead,
            q_dtype,
            kv_dtype,
            is_sparse=False,
            fast_mode=False,
            num_kv_splits=num_kv_splits,
            intra_batch_mode=True,
        )
        cached = tuple(torch.empty(shape, dtype=dtype, device=device) for shape, dtype in info)
        _WORKSPACE_CACHE[cache_key] = cached
    return cached


def _get_kv_indices(total_kv_len: int, device: torch.device) -> torch.Tensor:
    cache_key = (_device_cache_key(device), total_kv_len)
    cached = _KV_INDEX_CACHE.get(cache_key)
    if cached is None:
        cached = torch.arange(total_kv_len, dtype=torch.int32, device=device)
        _KV_INDEX_CACHE[cache_key] = cached
    return cached


def _get_uniform_indptrs(
    batch_size: int,
    q_seq_len: int,
    kv_seq_len: int,
    device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
    cache_key = (_device_cache_key(device), batch_size, q_seq_len, kv_seq_len)
    cached = _INDPTR_CACHE.get(cache_key)
    if cached is None:
        steps = torch.arange(batch_size + 1, dtype=torch.int32, device=device)
        cached = (
            steps * q_seq_len,
            steps * kv_seq_len,
            torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device),
        )
        _INDPTR_CACHE[cache_key] = cached
    return cached


def _get_output_buffer(shape: tuple[int, int, int], device: torch.device) -> torch.Tensor:
    cache_key = (_device_cache_key(device), shape)
    cached = _OUTPUT_CACHE.get(cache_key)
    if cached is None:
        cached = torch.empty(shape, dtype=torch.bfloat16, device=device)
        _OUTPUT_CACHE[cache_key] = cached
    return cached


def _select_num_kv_splits(batch_size: int, kv_seq_len: int) -> int:
    override = SPLIT_OVERRIDES.get((batch_size, kv_seq_len))
    if override is not None:
        return override
    total_kv = batch_size * kv_seq_len
    if total_kv >= 1_000_000:
        return 32
    if total_kv >= 131_072:
        return 24
    return 16


def make_mla_decode_metadata(
    batch_size: int,
    max_q_len: int,
    kv_seq_len: int,
    nhead: int,
    nhead_kv: int,
    q_dtype: torch.dtype,
    kv_dtype: torch.dtype,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    kv_last_page_len: torch.Tensor,
    num_kv_splits: int = NUM_KV_SPLITS,
):
    cache_key = (
        _device_cache_key(qo_indptr.device),
        batch_size,
        max_q_len,
        kv_seq_len,
        nhead,
        nhead_kv,
        q_dtype,
        kv_dtype,
        num_kv_splits,
    )
    cached = _METADATA_CACHE.get(cache_key)
    if cached is not None:
        return cached

    (
        work_metadata,
        work_indptr,
        work_info_set,
        reduce_indptr,
        reduce_final_map,
        reduce_partial_map,
    ) = _get_metadata_workspace(
        batch_size,
        max_q_len,
        nhead,
        nhead_kv,
        q_dtype,
        kv_dtype,
        num_kv_splits,
        qo_indptr.device,
    )

    get_mla_metadata_v1(
        qo_indptr,
        kv_indptr,
        kv_last_page_len,
        nhead // nhead_kv,
        nhead_kv,
        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, KV_GRANULARITY),
        max_seqlen_qo=max_q_len,
        uni_seqlen_qo=max_q_len,
        fast_mode=False,
        max_split_per_batch=num_kv_splits,
        intra_batch_mode=True,
        dtype_q=q_dtype,
        dtype_kv=kv_dtype,
    )

    cached = {
        "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,
    }
    _METADATA_CACHE[cache_key] = cached
    return cached


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

    batch_size = config["batch_size"]
    nq = config["num_heads"]
    nkv = config["num_kv_heads"]
    dq = config["qk_head_dim"]
    dv = config["v_head_dim"]
    q_seq_len = config["q_seq_len"]
    kv_seq_len = config["kv_seq_len"]

    qo_indptr, kv_indptr, kv_last_page_len = _get_uniform_indptrs(
        batch_size,
        q_seq_len,
        kv_seq_len,
        q.device,
    )

    q_fp8, q_scale = quantize_fp8(q)

    kv_buffer_fp8, kv_scale = kv_data["fp8"]
    total_kv_len = kv_buffer_fp8.shape[0]
    kv_indices = _get_kv_indices(total_kv_len, q.device)
    kv_buffer_4d = kv_buffer_fp8.reshape(total_kv_len, PAGE_SIZE, nkv, kv_buffer_fp8.shape[-1])

    num_kv_splits = _select_num_kv_splits(batch_size, kv_seq_len)
    meta = make_mla_decode_metadata(
        batch_size,
        q_seq_len,
        kv_seq_len,
        nq,
        nkv,
        q_fp8.dtype,
        kv_buffer_4d.dtype,
        qo_indptr,
        kv_indptr,
        kv_last_page_len,
        num_kv_splits=num_kv_splits,
    )

    o = _get_output_buffer((q.shape[0], nq, dv), q.device)
    mla_decode_fwd(
        q_fp8.reshape(-1, nq, dq),
        kv_buffer_4d,
        o,
        qo_indptr,
        kv_indptr,
        kv_indices,
        kv_last_page_len,
        q_seq_len,
        page_size=PAGE_SIZE,
        nhead_kv=nkv,
        sm_scale=SM_SCALE,
        logit_cap=0.0,
        num_kv_splits=num_kv_splits,
        q_scale=q_scale,
        kv_scale=kv_scale,
        intra_batch_mode=True,
        **meta,
    )
    return o
scrolls · 302 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