submission 709239
Maxwell Cipher · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 240 lines, June 9 Researcher Reciprocity License v1.0.
MLA_high_09.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-709239?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:b48c55ea9ed4f8886e43cf40bd5b013f6e4bf51c6ebb156524ecba90c931d4a0
license declaredunknown
license concludedunknown
authorsMaxwell Cipher
imported2026-08-15
Kernel source
MLA_high_09.py240 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
MLA policy sweep snapshot 09.
- 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
"""
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
SHORT_PAGE_SIZE = 1
LONG_PAGE_SIZE = 8
SHORT_Q_MODE = "fp8"
LONG_Q_MODE = "fp8"
TOKEN_PAGED_INDPTR = False
USE_LEGACY_SHORT_BASELINE = True
LEGACY_SHORT_MAX_BATCH = 32
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 _page_size_for_shape(kv_seq_len: int) -> int:
return SHORT_PAGE_SIZE if kv_seq_len <= 1024 else LONG_PAGE_SIZE
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
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 _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=max(page_size, 16),
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(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_BASELINE and kv_seq_len <= 1024 and batch_size <= LEGACY_SHORT_MAX_BATCH:
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 · 240 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 700955.
⋯ diff truncated: revisions differ almost entirely
Best evidence level for this revision: reported
JSON