Skip to content
KernelIndex
Search⌘K

submission 709818

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:1a23b9eea21144ecdf86e7f2e4b81589542ba30604026e8be047990b7e47bbc0
license declaredunknown
license concludedunknown
authorsMaxwell Cipher
imported2026-08-15

Kernel source

MLA_high_19.py273 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""
MLA policy sweep snapshot 19.

- 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
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): 16,
}

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


@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)
    block = 4096
    _cast_fp8_fixed[(triton.cdiv(flat_in.numel(), block),)](
        flat_in,
        flat_out,
        INV_Q_SCALE,
        flat_in.numel(),
        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)


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:
    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)
    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
    )
    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(
        page_size=page_size,
        nhead_kv=NUM_KV_HEADS,
        sm_scale=SM_SCALE,
        logit_cap=0.0,
        num_kv_splits=splits,
        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


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 _use_legacy_short(batch_size, kv_seq_len):
        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 709239.

⋯ 1 unchanged lines
#!POPCORN gpu MI355X
"""
- MLA policy sweep snapshot 09.
+ MLA policy sweep snapshot 19.
- - Keeps the legacy short-shape baseline for bs<=32, kv=1024
- - Uses the new fixed-range fp8 path elsewhere
- - Page size 1 for short paged routes, page size 8 for long routes
+ - 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
⋯ 17 unchanged lines
FIXED_Q_AMAX = 16.0
FIXED_Q_SCALE = FIXED_Q_AMAX / FP8_MAX
INV_Q_SCALE = FP8_MAX / FIXED_Q_AMAX
- SHORT_PAGE_SIZE = 1
- LONG_PAGE_SIZE = 8
- SHORT_Q_MODE = "fp8"
- LONG_Q_MODE = "fp8"
+ 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_SHORT_BASELINE = True
- LEGACY_SHORT_MAX_BATCH = 32
+ 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,
⋯ 24 unchanged lines
return -1 if device.index is None else int(device.index)
- def _page_size_for_shape(kv_seq_len: int) -> int:
- return SHORT_PAGE_SIZE if kv_seq_len <= 1024 else LONG_PAGE_SIZE
+ 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:
- del batch_size
- return SHORT_Q_MODE if kv_seq_len <= 1024 else LONG_Q_MODE
+ 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)]
⋯ 2 unchanged lines
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
⋯ 40 unchanged lines
rfm,
rpm,
page_size=page_size,
- kv_granularity=max(page_size, 16),
+ 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,
⋯ 65 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(kv_seq_len)
+ 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)
⋯ 24 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_BASELINE and kv_seq_len <= 1024 and batch_size <= LEGACY_SHORT_MAX_BATCH:
+ if _use_legacy_short(batch_size, kv_seq_len):
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 · 123 diff lines total

Best evidence level for this revision: reported

JSON