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
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 torchimport tritonimport triton.language as tl⋯ 9 unchanged linesSM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)FP8_DTYPE = aiter_dtypes.fp8FP8_MAX = float(torch.finfo(FP8_DTYPE).max)-FIXED_Q_AMAX = 16.0FIXED_Q_SCALE = FIXED_Q_AMAX / FP8_MAXINV_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 linesdef _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 linesdef _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 linesq, kv_data, qo_indptr, kv_indptr, config = databatch_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