submission 739084
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_5.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-739084?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
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:4ca2c8593f6224bc79563cfe9feb608e7c0ca65f0003cd2c22ccbbd643f2e8c4
license declaredunknown
license concludedunknown
authorsMaxwell Cipher
imported2026-08-15
Kernel source
MLA_5.py273 lines
#!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
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]] = {}
@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 709818.
⋯ 1 unchanged lines#!POPCORN gpu MI355X"""- MLA policy sweep snapshot 19.+ 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⋯ 44 unchanged lines(64, 1024): 8,(64, 8192): 16,(256, 1024): 16,- (256, 8192): 16,+ (256, 8192): 32,}_META_CACHE: dict[tuple, tuple] = {}
Best evidence level for this revision: reported
JSON