Skip to content
KernelIndex
Search⌘K

submission 741445

Maxwell Cipher · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9df3c97762f6d0a250eba4f2d9235a8a663d613a015396a1a9a68dc1aab536d1
license declaredunknown
license concludedunknown
authorsMaxwell Cipher
imported2026-08-15

Kernel source

MLA_v5d.py273 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
from __future__ import annotations
import torch
import triton
import triton.language as tl
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

NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
FP8_DTYPE = aiter_dtypes.fp8
FP8_MAX = float(torch.finfo(FP8_DTYPE).max)
FIXED_Q_AMAX = 16.0
FIXED_Q_SCALE = FIXED_Q_AMAX / FP8_MAX
INV_Q_SCALE = FP8_MAX / FIXED_Q_AMAX
DEFAULT_SHORT_PAGE_SIZE = 1
DEFAULT_LONG_PAGE_SIZE = 8
DEFAULT_SHORT_Q_MODE = "fp8"
DEFAULT_LONG_Q_MODE = "fp8"
PAGE_SIZE_OVERRIDES: dict[tuple[int, int], int] = {
    (32, 1024): 2,
    (64, 1024): 2,
    (256, 1024): 2,
}
Q_MODE_OVERRIDES: dict[tuple[int, int], str] = {}
TOKEN_PAGED_INDPTR = False
USE_LEGACY_SHAPES: frozenset[tuple[int, int]] = frozenset({(4, 1024)})
KV_GRANULARITY_MODE = "inverse_page"
KV_GRANULARITY_OVERRIDES: dict[tuple[int, int], int] = {}
FAST_MODE_SHAPES: frozenset[tuple[int, int]] = frozenset()
SPLIT_PLAN = {
    (4, 1024): 4,
    (4, 8192): 16,
    (32, 1024): 8,
    (32, 8192): 16,
    (64, 1024): 8,
    (64, 8192): 16,
    (256, 1024): 16,
    (256, 8192): 32,
}

_META_CACHE: dict[tuple, tuple] = {}
_LEGACY_PAGE_CACHE: dict[tuple, tuple[torch.Tensor, torch.Tensor]] = {}
_Q_SCALE_CACHE: dict[int, torch.Tensor] = {}
_GRID_CACHE: dict[int, tuple] = {}


@triton.jit
def _cast_fp8_fixed(src_ptr, dst_ptr, inv_scale, n_elements, FP8_LIMIT: tl.constexpr, BLOCK: tl.constexpr):
    pid = tl.program_id(0)
    offs = pid * BLOCK + tl.arange(0, BLOCK)
    mask = offs < n_elements
    x = tl.load(src_ptr + offs, mask=mask, other=0.0).to(tl.float32)
    y = tl.clamp(x * inv_scale, -FP8_LIMIT, FP8_LIMIT)
    tl.store(dst_ptr + offs, y.to(dst_ptr.dtype.element_ty), mask=mask)


def _device_key(device: torch.device) -> int:
    return -1 if device.index is None else int(device.index)


def _shape_key(batch_size: int, kv_seq_len: int) -> tuple[int, int]:
    return (batch_size, kv_seq_len)


def _q_mode_for_shape(batch_size: int, kv_seq_len: int) -> str:
    return Q_MODE_OVERRIDES.get(
        _shape_key(batch_size, kv_seq_len),
        DEFAULT_SHORT_Q_MODE if kv_seq_len <= 1024 else DEFAULT_LONG_Q_MODE,
    )


def _page_size_for_shape(batch_size: int, kv_seq_len: int) -> int:
    return PAGE_SIZE_OVERRIDES.get(
        _shape_key(batch_size, kv_seq_len),
        DEFAULT_SHORT_PAGE_SIZE if kv_seq_len <= 1024 else DEFAULT_LONG_PAGE_SIZE,
    )


def _split_for_shape(batch_size: int, kv_seq_len: int) -> int:
    return SPLIT_PLAN[(batch_size, kv_seq_len)]


def _fast_mode_for_shape(batch_size: int, kv_seq_len: int) -> bool:
    return (batch_size, kv_seq_len) in FAST_MODE_SHAPES


def _use_legacy_short(batch_size: int, kv_seq_len: int) -> bool:
    return (batch_size, kv_seq_len) in USE_LEGACY_SHAPES


def _kv_granularity_for_page(page_size: int) -> int:
    if KV_GRANULARITY_MODE == "inverse_page":
        return max(1, 16 // page_size)
    return max(page_size, 16)


def _kv_granularity_for_shape(batch_size: int, kv_seq_len: int, page_size: int) -> int:
    return KV_GRANULARITY_OVERRIDES.get(
        _shape_key(batch_size, kv_seq_len),
        _kv_granularity_for_page(page_size),
    )


def _shape_indptr(batch_size: int, kv_seq_len: int, page_size: int, token_mode: bool, device: torch.device) -> tuple[torch.Tensor, torch.Tensor, int]:
    if page_size == 1 or token_mode:
        kv_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device) * kv_seq_len
    else:
        kv_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device) * (kv_seq_len // page_size)
    last_page = kv_seq_len if page_size == 1 else page_size
    kv_last_page_len = torch.full((batch_size,), last_page, dtype=torch.int32, device=device)
    num_pages = batch_size * (kv_seq_len if page_size == 1 else kv_seq_len // page_size)
    return kv_indptr, kv_last_page_len, num_pages


def _get_meta(batch_size: int, kv_seq_len: int, page_size: int, q_mode: str, splits: int, fast_mode: bool, device: torch.device) -> tuple[dict[str, torch.Tensor], torch.Tensor, torch.Tensor, torch.Tensor]:
    key = (_device_key(device), batch_size, kv_seq_len, page_size, q_mode, splits, fast_mode, TOKEN_PAGED_INDPTR)
    if key not in _META_CACHE:
        q_dtype = FP8_DTYPE if q_mode == "fp8" else torch.bfloat16
        kv_indptr_arg, kv_last_page_len, num_pages = _shape_indptr(
            batch_size, kv_seq_len, page_size, TOKEN_PAGED_INDPTR, device
        )
        info = get_mla_metadata_info_v1(
            batch_size,
            1,
            NUM_HEADS,
            q_dtype,
            FP8_DTYPE,
            is_sparse=False,
            fast_mode=fast_mode,
            num_kv_splits=splits,
            intra_batch_mode=True,
        )
        work = [torch.empty(shape, dtype=dtype, device=device) for shape, dtype in info]
        wm, wi, wis, ri, rfm, rpm = work
        qo_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device)
        get_mla_metadata_v1(
            qo_indptr,
            kv_indptr_arg,
            kv_last_page_len,
            NUM_HEADS // NUM_KV_HEADS,
            NUM_KV_HEADS,
            True,
            wm,
            wis,
            wi,
            ri,
            rfm,
            rpm,
            page_size=page_size,
            kv_granularity=_kv_granularity_for_shape(batch_size, kv_seq_len, page_size),
            max_seqlen_qo=1,
            uni_seqlen_qo=1,
            fast_mode=fast_mode,
            max_split_per_batch=splits,
            intra_batch_mode=True,
            dtype_q=q_dtype,
            dtype_kv=FP8_DTYPE,
        )
        meta = {
            "work_meta_data": wm,
            "work_indptr": wi,
            "work_info_set": wis,
            "reduce_indptr": ri,
            "reduce_final_map": rfm,
            "reduce_partial_map": rpm,
        }
        kv_indices = torch.arange(num_pages, dtype=torch.int32, device=device)
        _META_CACHE[key] = (meta, kv_indices, kv_indptr_arg, kv_last_page_len)
    return _META_CACHE[key]


def _get_legacy_pg1(batch_size: int, kv_seq_len: int, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]:
    key = (_device_key(device), batch_size, kv_seq_len)
    if key not in _LEGACY_PAGE_CACHE:
        _LEGACY_PAGE_CACHE[key] = (
            torch.arange(batch_size * kv_seq_len, dtype=torch.int32, device=device),
            torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device),
        )
    return _LEGACY_PAGE_CACHE[key]


def _alloc_out(q: torch.Tensor) -> torch.Tensor:
    return torch.empty((q.shape[0], NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=q.device)


def _quantize_q_fp8(q: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    flat_in = q.reshape(-1)
    flat_out = torch.empty((flat_in.numel(),), dtype=FP8_DTYPE, device=q.device)
    n = flat_in.numel()
    block = 4096
    if n not in _GRID_CACHE:
        _GRID_CACHE[n] = (triton.cdiv(n, block),)
    _cast_fp8_fixed[_GRID_CACHE[n]](
        flat_in,
        flat_out,
        INV_Q_SCALE,
        n,
        FP8_LIMIT=FP8_MAX,
        BLOCK=block,
    )
    dk = -1 if q.device.index is None else int(q.device.index)
    if dk not in _Q_SCALE_CACHE:
        _Q_SCALE_CACHE[dk] = torch.full((1,), FIXED_Q_SCALE, dtype=torch.float32, device=q.device)
    return flat_out.view(q.shape[0], NUM_HEADS, QK_HEAD_DIM), _Q_SCALE_CACHE[dk]


def _run_legacy_short(q: torch.Tensor, kv_bf16: torch.Tensor, qo_indptr: torch.Tensor, kv_indptr: torch.Tensor, batch_size: int, kv_seq_len: int) -> torch.Tensor:
    out = _alloc_out(q)
    page_ids, last_page_len = _get_legacy_pg1(batch_size, kv_seq_len, q.device)
    mla_decode_fwd(
        q.view(-1, NUM_HEADS, QK_HEAD_DIM),
        kv_bf16.view(-1, 1, NUM_KV_HEADS, kv_bf16.shape[-1]),
        out,
        qo_indptr,
        kv_indptr,
        page_ids,
        last_page_len,
        1,
        page_size=1,
        nhead_kv=NUM_KV_HEADS,
        sm_scale=SM_SCALE,
        intra_batch_mode=False,
    )
    return out


def _run_paged(q: torch.Tensor, kv_fp8: torch.Tensor, kv_scale: torch.Tensor, qo_indptr: torch.Tensor, batch_size: int, kv_seq_len: int) -> torch.Tensor:
    sk = (batch_size, kv_seq_len)
    page_size = PAGE_SIZE_OVERRIDES.get(sk, DEFAULT_SHORT_PAGE_SIZE if kv_seq_len <= 1024 else DEFAULT_LONG_PAGE_SIZE)
    splits = SPLIT_PLAN[sk]
    meta, kv_indices, kv_indptr_arg, kv_last_page_len = _get_meta(
        batch_size, kv_seq_len, page_size, "fp8", splits, False, q.device
    )
    out = _alloc_out(q)
    q_view, q_scale = _quantize_q_fp8(q.view(-1, NUM_HEADS, QK_HEAD_DIM))
    mla_decode_fwd(
        q_view,
        kv_fp8.view(-1, page_size, NUM_KV_HEADS, kv_fp8.shape[-1]),
        out,
        qo_indptr,
        kv_indptr_arg,
        kv_indices,
        kv_last_page_len,
        1,
        page_size=page_size,
        nhead_kv=NUM_KV_HEADS,
        sm_scale=SM_SCALE,
        logit_cap=0.0,
        num_kv_splits=splits,
        q_scale=q_scale,
        kv_scale=kv_scale,
        intra_batch_mode=True,
        **meta,
    )
    return out


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    batch_size = int(config["batch_size"])
    kv_seq_len = int(config["kv_seq_len"])
    if (batch_size, kv_seq_len) in USE_LEGACY_SHAPES:
        return _run_legacy_short(q, kv_data["bf16"], qo_indptr, kv_indptr, batch_size, kv_seq_len)
    kv_fp8, kv_scale = kv_data["fp8"]
    return _run_paged(q, kv_fp8, kv_scale, qo_indptr, batch_size, kv_seq_len)
scrolls · 273 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 739084.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
-
- """
- Variant 5: splits=32 for (256,8192)
-
- - Keeps only the lightest 1k corner on the legacy short baseline
- - Probes page size 2 on the three heavier 1k batches
- - Uses inverse-page kv granularity on paged metadata
- """
-
from __future__ import annotations
-
import torch
import triton
import triton.language as tl
⋯ 9 unchanged lines
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
FP8_DTYPE = aiter_dtypes.fp8
FP8_MAX = float(torch.finfo(FP8_DTYPE).max)
-
FIXED_Q_AMAX = 16.0
FIXED_Q_SCALE = FIXED_Q_AMAX / FP8_MAX
INV_Q_SCALE = FP8_MAX / FIXED_Q_AMAX
⋯ 25 unchanged lines
_META_CACHE: dict[tuple, tuple] = {}
_LEGACY_PAGE_CACHE: dict[tuple, tuple[torch.Tensor, torch.Tensor]] = {}
+ _Q_SCALE_CACHE: dict[int, torch.Tensor] = {}
+ _GRID_CACHE: dict[int, tuple] = {}
@triton.jit
⋯ 138 unchanged lines
def _quantize_q_fp8(q: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
flat_in = q.reshape(-1)
flat_out = torch.empty((flat_in.numel(),), dtype=FP8_DTYPE, device=q.device)
+ n = flat_in.numel()
block = 4096
- _cast_fp8_fixed[(triton.cdiv(flat_in.numel(), block),)](
+ if n not in _GRID_CACHE:
+ _GRID_CACHE[n] = (triton.cdiv(n, block),)
+ _cast_fp8_fixed[_GRID_CACHE[n]](
flat_in,
flat_out,
INV_Q_SCALE,
- flat_in.numel(),
+ n,
FP8_LIMIT=FP8_MAX,
BLOCK=block,
)
- return flat_out.view(q.shape[0], NUM_HEADS, QK_HEAD_DIM), torch.full((1,), FIXED_Q_SCALE, dtype=torch.float32, device=q.device)
+ dk = -1 if q.device.index is None else int(q.device.index)
+ if dk not in _Q_SCALE_CACHE:
+ _Q_SCALE_CACHE[dk] = torch.full((1,), FIXED_Q_SCALE, dtype=torch.float32, device=q.device)
+ return flat_out.view(q.shape[0], NUM_HEADS, QK_HEAD_DIM), _Q_SCALE_CACHE[dk]
def _run_legacy_short(q: torch.Tensor, kv_bf16: torch.Tensor, qo_indptr: torch.Tensor, kv_indptr: torch.Tensor, batch_size: int, kv_seq_len: int) -> torch.Tensor:
⋯ 17 unchanged lines
def _run_paged(q: torch.Tensor, kv_fp8: torch.Tensor, kv_scale: torch.Tensor, qo_indptr: torch.Tensor, batch_size: int, kv_seq_len: int) -> torch.Tensor:
- page_size = _page_size_for_shape(batch_size, kv_seq_len)
- q_mode = _q_mode_for_shape(batch_size, kv_seq_len)
- splits = _split_for_shape(batch_size, kv_seq_len)
- fast_mode = _fast_mode_for_shape(batch_size, kv_seq_len)
+ sk = (batch_size, kv_seq_len)
+ page_size = PAGE_SIZE_OVERRIDES.get(sk, DEFAULT_SHORT_PAGE_SIZE if kv_seq_len <= 1024 else DEFAULT_LONG_PAGE_SIZE)
+ splits = SPLIT_PLAN[sk]
meta, kv_indices, kv_indptr_arg, kv_last_page_len = _get_meta(
- batch_size, kv_seq_len, page_size, q_mode, splits, fast_mode, q.device
+ batch_size, kv_seq_len, page_size, "fp8", splits, False, q.device
)
out = _alloc_out(q)
- kv_view = kv_fp8.view(-1, page_size, NUM_KV_HEADS, kv_fp8.shape[-1])
- q_view = q.view(-1, NUM_HEADS, QK_HEAD_DIM)
- call_kwargs = dict(
+ q_view, q_scale = _quantize_q_fp8(q.view(-1, NUM_HEADS, QK_HEAD_DIM))
+ mla_decode_fwd(
+ q_view,
+ kv_fp8.view(-1, page_size, NUM_KV_HEADS, kv_fp8.shape[-1]),
+ out,
+ qo_indptr,
+ kv_indptr_arg,
+ kv_indices,
+ kv_last_page_len,
+ 1,
page_size=page_size,
nhead_kv=NUM_KV_HEADS,
sm_scale=SM_SCALE,
logit_cap=0.0,
num_kv_splits=splits,
+ q_scale=q_scale,
kv_scale=kv_scale,
intra_batch_mode=True,
**meta,
)
- if q_mode == "fp8":
- q_view, q_scale = _quantize_q_fp8(q_view)
- call_kwargs["q_scale"] = q_scale
- mla_decode_fwd(q_view, kv_view, out, qo_indptr, kv_indptr_arg, kv_indices, kv_last_page_len, 1, **call_kwargs)
return out
⋯ 1 unchanged lines
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = int(config["batch_size"])
kv_seq_len = int(config["kv_seq_len"])
- if _use_legacy_short(batch_size, kv_seq_len):
+ if (batch_size, kv_seq_len) in USE_LEGACY_SHAPES:
return _run_legacy_short(q, kv_data["bf16"], qo_indptr, kv_indptr, batch_size, kv_seq_len)
kv_fp8, kv_scale = kv_data["fp8"]
return _run_paged(q, kv_fp8, kv_scale, qo_indptr, batch_size, kv_seq_len)
scrolls · 114 diff lines total

Best evidence level for this revision: reported

JSON