submission 723086
.jonnss · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 589 lines, June 9 Researcher Reciprocity License v1.0.
Submission.v208.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-723086?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:061d3e148ef19bf616b5fb60d635de3e59cb4f79d5c0dd571ad2c5fd53951cb3
license declaredunknown
license concludedunknown
authors.jonnss
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
_NON_PERSISTENT_SHAPES: set[tuple[int, int]] = set()Kernel source
Submission.v208.py589 lines
# gpumode leaderboard reference
"""
MLA-SESSION true b4 hybrid.
Keep the live `v158` runner unchanged for every non-`b4` shape.
Only the two `b4` shapes get the old exact-shape-bank runner family:
- `(4, 1024)` uses the `v167` exact-shape runner depth
- `(4, 8192)` uses the `v173` deeper-Q-ring exact-shape runner depth
This targets the only remaining artifact-backed ranked gap after the `b256`
hybrid lane died in `v206` and `v207`.
"""
import os
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
import torch
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
Q_DTYPE = os.getenv("AMD2_MLA_Q_DTYPE", "fp8").lower()
KV_DTYPE = os.getenv("AMD2_MLA_KV_DTYPE", "fp8").lower()
if Q_DTYPE not in {"fp8", "bf16"}:
raise ValueError(f"Unsupported AMD2_MLA_Q_DTYPE={Q_DTYPE!r}")
if KV_DTYPE not in {"fp8", "bf16"}:
raise ValueError(f"Unsupported AMD2_MLA_KV_DTYPE={KV_DTYPE!r}")
_NON_PERSISTENT_SHAPES: set[tuple[int, int]] = set()
_SHAPE_CONFIG: dict[tuple[int, int], tuple[int, int, int, bool]] = {
(4, 1024): (1, 8, 128, False),
(4, 8192): (8, 32, 128, False),
(32, 1024): (1, 4, 128, False),
(32, 8192): (8, 32, 32, True),
(64, 1024): (2, 8, 128, False),
(64, 8192): (8, 32, 32, False),
(256, 1024): (2, 8, 32, False),
(256, 8192): (8, 32, 32, False),
}
_Q_CACHE_SLOTS = 16
_OUTPUT_CACHE_SLOTS = 4
_FAST_PATH_SHAPES = {
(4, 1024),
(4, 8192),
}
_EXACT_Q_CACHE_SLOTS_BY_SHAPE = {
(4, 1024): 16,
(4, 8192): 32,
}
_EXACT_SAFE_Q_CACHE_SLOTS = 16
torch.set_grad_enabled(False)
_KV_INDICES_CACHE = {}
_PERSISTENT_SETUP_CACHE = {}
_NON_PERSISTENT_SETUP_CACHE = {}
_PERSISTENT_RUNNER_CACHE = {}
_EXACT_SHAPE_RUNNER_BANK = {}
_UNIT_Q_SCALE_CACHE = {}
_EXACT_SHAPE_CONFIGS = {
shape: {
"batch_size": shape[0],
"q_seq_len": 1,
"kv_seq_len": shape[1],
"num_heads": NUM_HEADS,
"num_kv_heads": NUM_KV_HEADS,
"qk_head_dim": QK_HEAD_DIM,
"v_head_dim": V_HEAD_DIM,
"sm_scale": SM_SCALE,
}
for shape in _FAST_PATH_SHAPES
}
def _get_unit_q_scale(device: torch.device) -> torch.Tensor:
key = device.index or 0
cached = _UNIT_Q_SCALE_CACHE.get(key)
if cached is None:
cached = torch.ones((1,), dtype=torch.float32, device=device)
_UNIT_Q_SCALE_CACHE[key] = cached
return cached
def _quantize_fp8_copy(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
out = torch.empty(tensor.shape, dtype=FP8_DTYPE, device=tensor.device)
out.copy_(tensor)
return out, _get_unit_q_scale(tensor.device)
def _get_shape_config(config: dict) -> tuple[int, int, int, bool]:
batch_size = int(config["batch_size"])
kv_seq_len = int(config["kv_seq_len"])
shape = (batch_size, kv_seq_len)
exact = _SHAPE_CONFIG.get(shape)
if exact is not None:
return exact
if kv_seq_len >= 8192:
if batch_size <= 4:
return 8, 32, 128, False
if batch_size <= 32:
return 8, 32, 32, True
return 8, 32, 32, False
if batch_size <= 32:
return 1, 8, 128, False
if batch_size <= 64:
return 2, 8, 128, False
return 2, 8, 32, False
def _is_non_persistent_case(config: dict) -> bool:
return (int(config["batch_size"]), int(config["kv_seq_len"])) in _NON_PERSISTENT_SHAPES
def _get_kv_indices(total_pages: int, device: torch.device) -> torch.Tensor:
key = (device.index or 0, total_pages)
cached = _KV_INDICES_CACHE.get(key)
if cached is None:
cached = torch.arange(total_pages, dtype=torch.int32, device=device)
_KV_INDICES_CACHE[key] = cached
return cached
def _make_metadata(
qo_indptr,
kv_indptr,
kv_last_page_len,
nhead,
nhead_kv,
q_dtype,
kv_dtype,
page_size,
num_kv_splits,
kv_granularity,
intra_batch_mode,
):
batch_size = qo_indptr.numel() - 1
max_q_len = int((qo_indptr[1] - qo_indptr[0]).item())
info = get_mla_metadata_info_v1(
batch_size, max_q_len, nhead, q_dtype, kv_dtype,
is_sparse=False, fast_mode=False, num_kv_splits=num_kv_splits, intra_batch_mode=intra_batch_mode,
)
work = [torch.empty(shape, dtype=dtype, device=qo_indptr.device) for shape, dtype in info]
work_metadata, work_indptr, work_info_set, reduce_indptr, reduce_final_map, reduce_partial_map = work
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_last_page_len,
nhead // nhead_kv, nhead_kv, True,
work_metadata, work_info_set, work_indptr,
reduce_indptr, reduce_final_map, reduce_partial_map,
page_size=page_size, kv_granularity=kv_granularity,
max_seqlen_qo=max_q_len, uni_seqlen_qo=max_q_len,
fast_mode=False, max_split_per_batch=num_kv_splits, intra_batch_mode=intra_batch_mode,
dtype_q=q_dtype, dtype_kv=kv_dtype,
)
return {
"work_meta_data": work_metadata,
"work_indptr": work_indptr,
"work_info_set": work_info_set,
"reduce_indptr": reduce_indptr,
"reduce_final_map": reduce_final_map,
"reduce_partial_map": reduce_partial_map,
}
def _get_non_persistent_setup(batch_size: int, q_seq_len: int, kv_seq_len: int, device: torch.device):
key = (device.index or 0, batch_size, q_seq_len, kv_seq_len)
cached = _NON_PERSISTENT_SETUP_CACHE.get(key)
if cached is None:
qo_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device) * q_seq_len
kv_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device) * kv_seq_len
cached = {
"qo_indptr": qo_indptr,
"kv_indptr": kv_indptr,
"kv_indices": _get_kv_indices(batch_size * kv_seq_len, device),
"kv_last_page_len": torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device),
}
_NON_PERSISTENT_SETUP_CACHE[key] = cached
return cached
def _get_persistent_setup(
config: dict,
q_dtype: torch.dtype,
kv_dtype: torch.dtype,
page_size: int,
num_kv_splits: int,
kv_granularity: int,
intra_batch_mode: bool,
device: torch.device,
):
batch_size = int(config["batch_size"])
q_seq_len = int(config["q_seq_len"])
kv_seq_len = int(config["kv_seq_len"])
nhead = int(config["num_heads"])
nhead_kv = int(config["num_kv_heads"])
pages_per_seq = (kv_seq_len + page_size - 1) // page_size
key = (
device.index or 0,
batch_size,
q_seq_len,
kv_seq_len,
nhead,
nhead_kv,
q_dtype,
kv_dtype,
page_size,
num_kv_splits,
kv_granularity,
intra_batch_mode,
)
cached = _PERSISTENT_SETUP_CACHE.get(key)
if cached is None:
qo_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device) * q_seq_len
kv_page_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device) * pages_per_seq
kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)
cached = {
"qo_indptr": qo_indptr,
"kv_page_indptr": kv_page_indptr,
"kv_page_indices": _get_kv_indices(batch_size * pages_per_seq, device),
"kv_last_page_len": kv_last_page_len,
"meta": _make_metadata(
qo_indptr,
kv_page_indptr,
kv_last_page_len,
nhead,
nhead_kv,
q_dtype,
kv_dtype,
page_size,
num_kv_splits,
kv_granularity,
intra_batch_mode,
),
}
_PERSISTENT_SETUP_CACHE[key] = cached
return cached
def _reshape_paged_kv(kv_buffer: torch.Tensor, batch_size: int, kv_seq_len: int, page_size: int, nhead_kv: int, qk_head_dim: int) -> torch.Tensor:
total_kv = batch_size * kv_seq_len
if kv_buffer.shape[0] != total_kv:
raise ValueError(f"Expected uniform total_kv={total_kv}, got {kv_buffer.shape[0]}")
if kv_seq_len % page_size != 0:
pages_per_seq = (kv_seq_len + page_size - 1) // page_size
kv_paged = torch.zeros((batch_size * pages_per_seq, page_size, nhead_kv, qk_head_dim), dtype=kv_buffer.dtype, device=kv_buffer.device)
kv_seq = kv_buffer.reshape(batch_size, kv_seq_len, nhead_kv, qk_head_dim)
for batch_idx in range(batch_size):
flat = kv_paged[batch_idx * pages_per_seq:(batch_idx + 1) * pages_per_seq].view(pages_per_seq * page_size, nhead_kv, qk_head_dim)
flat[:kv_seq_len].copy_(kv_seq[batch_idx])
return kv_paged
pages_per_seq = kv_seq_len // page_size
return kv_buffer.reshape(batch_size, kv_seq_len, nhead_kv, qk_head_dim).reshape(batch_size * pages_per_seq, page_size, nhead_kv, qk_head_dim)
def _run_non_persistent(q, kv_buffer, config, q_scale, kv_scale):
batch_size = int(config["batch_size"])
q_seq_len = int(config["q_seq_len"])
kv_seq_len = int(config["kv_seq_len"])
nhead = int(config["num_heads"])
nhead_kv = int(config["num_kv_heads"])
qk_head_dim = int(config["qk_head_dim"])
v_head_dim = int(config["v_head_dim"])
setup = _get_non_persistent_setup(batch_size, q_seq_len, kv_seq_len, q.device)
out = _get_output_workspace(q.shape[0], nhead, v_head_dim, q.device)
mla_decode_fwd(
q.reshape(-1, nhead, qk_head_dim),
kv_buffer.reshape(kv_buffer.shape[0], 1, nhead_kv, qk_head_dim),
out,
setup["qo_indptr"], setup["kv_indptr"], setup["kv_indices"], setup["kv_last_page_len"], q_seq_len,
page_size=1, nhead_kv=nhead_kv, sm_scale=float(config.get("sm_scale", SM_SCALE)), logit_cap=0.0,
q_scale=q_scale, kv_scale=kv_scale,
)
return out
def _run_persistent(q, kv_buffer, config, q_scale, kv_scale):
batch_size = int(config["batch_size"])
kv_seq_len = int(config["kv_seq_len"])
nhead = int(config["num_heads"])
nhead_kv = int(config["num_kv_heads"])
qk_head_dim = int(config["qk_head_dim"])
v_head_dim = int(config["v_head_dim"])
q_seq_len = int(config["q_seq_len"])
page_size, num_kv_splits, kv_granularity, intra_batch_mode = _get_shape_config(config)
setup = _get_persistent_setup(
config,
q.dtype,
kv_buffer.dtype,
page_size,
num_kv_splits,
kv_granularity,
intra_batch_mode,
q.device,
)
out = _get_output_workspace(q.shape[0], nhead, v_head_dim, q.device)
mla_decode_fwd(
q.reshape(-1, nhead, qk_head_dim),
_reshape_paged_kv(kv_buffer, batch_size, kv_seq_len, page_size, nhead_kv, qk_head_dim),
out,
setup["qo_indptr"], setup["kv_page_indptr"], setup["kv_page_indices"], setup["kv_last_page_len"], q_seq_len,
page_size=page_size, nhead_kv=nhead_kv, sm_scale=float(config.get("sm_scale", SM_SCALE)), logit_cap=0.0,
num_kv_splits=num_kv_splits,
q_scale=q_scale,
kv_scale=kv_scale,
intra_batch_mode=intra_batch_mode,
**setup["meta"],
)
return out
def _make_persistent_runner(config: dict, kv_dtype: torch.dtype, device: torch.device):
batch_size = int(config["batch_size"])
q_seq_len = int(config["q_seq_len"])
kv_seq_len = int(config["kv_seq_len"])
nhead = int(config["num_heads"])
nhead_kv = int(config["num_kv_heads"])
qk_head_dim = int(config["qk_head_dim"])
v_head_dim = int(config["v_head_dim"])
num_tokens = batch_size * q_seq_len
page_size, num_kv_splits, kv_granularity, intra_batch_mode = _get_shape_config(config)
q_storage_dtype = FP8_DTYPE if Q_DTYPE == "fp8" else torch.bfloat16
setup = _get_persistent_setup(
config,
q_storage_dtype,
kv_dtype,
page_size,
num_kv_splits,
kv_granularity,
intra_batch_mode,
device,
)
qo_indptr = setup["qo_indptr"]
kv_page_indptr = setup["kv_page_indptr"]
kv_page_indices = setup["kv_page_indices"]
kv_last_page_len = setup["kv_last_page_len"]
meta = setup["meta"]
work_meta_data = meta["work_meta_data"]
work_indptr = meta["work_indptr"]
work_info_set = meta["work_info_set"]
reduce_indptr = meta["reduce_indptr"]
reduce_final_map = meta["reduce_final_map"]
reduce_partial_map = meta["reduce_partial_map"]
sm_scale = float(config.get("sm_scale", SM_SCALE))
pages_per_seq = (kv_seq_len + page_size - 1) // page_size
output_buffers = [
torch.empty((num_tokens, nhead, v_head_dim), dtype=torch.bfloat16, device=device)
for _ in range(_OUTPUT_CACHE_SLOTS)
]
output_index = 0
if Q_DTYPE == "fp8":
q_buffers = [
torch.empty((num_tokens, nhead, qk_head_dim), dtype=FP8_DTYPE, device=device)
for _ in range(_Q_CACHE_SLOTS)
]
q_index = 0
unit_q_scale = _get_unit_q_scale(device)
else:
q_buffers = []
q_index = 0
unit_q_scale = None
if kv_seq_len % page_size == 0:
def reshape_kv(kv_buffer: torch.Tensor) -> torch.Tensor:
return kv_buffer.reshape(batch_size * pages_per_seq, page_size, nhead_kv, qk_head_dim)
else:
def reshape_kv(kv_buffer: torch.Tensor) -> torch.Tensor:
return _reshape_paged_kv(kv_buffer, batch_size, kv_seq_len, page_size, nhead_kv, qk_head_dim)
def runner(q: torch.Tensor, kv_buffer: torch.Tensor, kv_scale: torch.Tensor | None) -> torch.Tensor:
nonlocal output_index, q_index
if Q_DTYPE == "fp8":
q_input = q_buffers[q_index]
q_input.copy_(q)
q_index = (q_index + 1) % len(q_buffers)
q_scale = unit_q_scale
else:
q_input = q
q_scale = None
out = output_buffers[output_index]
output_index = (output_index + 1) % len(output_buffers)
mla_decode_fwd(
q_input.reshape(-1, nhead, qk_head_dim),
reshape_kv(kv_buffer),
out,
qo_indptr,
kv_page_indptr,
kv_page_indices,
kv_last_page_len,
q_seq_len,
page_size=page_size,
nhead_kv=nhead_kv,
sm_scale=sm_scale,
logit_cap=0.0,
num_kv_splits=num_kv_splits,
q_scale=q_scale,
kv_scale=kv_scale,
intra_batch_mode=intra_batch_mode,
work_meta_data=work_meta_data,
work_indptr=work_indptr,
work_info_set=work_info_set,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
)
return out
return runner
def _get_persistent_runner(config: dict, kv_dtype: torch.dtype, device: torch.device):
key = (
device.index or 0,
int(config["batch_size"]),
int(config["q_seq_len"]),
int(config["kv_seq_len"]),
int(config["num_heads"]),
int(config["num_kv_heads"]),
int(config["qk_head_dim"]),
int(config["v_head_dim"]),
FP8_DTYPE if Q_DTYPE == "fp8" else torch.bfloat16,
kv_dtype,
)
cached = _PERSISTENT_RUNNER_CACHE.get(key)
if cached is None:
cached = _make_persistent_runner(config, kv_dtype, device)
_PERSISTENT_RUNNER_CACHE[key] = cached
return cached
def _make_exact_shape_runner(shape: tuple[int, int], kv_dtype: torch.dtype, device: torch.device):
config = _EXACT_SHAPE_CONFIGS[shape]
batch_size = int(config["batch_size"])
q_seq_len = int(config["q_seq_len"])
kv_seq_len = int(config["kv_seq_len"])
nhead = int(config["num_heads"])
nhead_kv = int(config["num_kv_heads"])
qk_head_dim = int(config["qk_head_dim"])
v_head_dim = int(config["v_head_dim"])
num_tokens = batch_size * q_seq_len
page_size, num_kv_splits, kv_granularity, intra_batch_mode = _get_shape_config(config)
q_storage_dtype = FP8_DTYPE if Q_DTYPE == "fp8" else torch.bfloat16
setup = _get_persistent_setup(
config,
q_storage_dtype,
kv_dtype,
page_size,
num_kv_splits,
kv_granularity,
intra_batch_mode,
device,
)
qo_indptr = setup["qo_indptr"]
kv_page_indptr = setup["kv_page_indptr"]
kv_page_indices = setup["kv_page_indices"]
kv_last_page_len = setup["kv_last_page_len"]
meta = setup["meta"]
work_meta_data = meta["work_meta_data"]
work_indptr = meta["work_indptr"]
work_info_set = meta["work_info_set"]
reduce_indptr = meta["reduce_indptr"]
reduce_final_map = meta["reduce_final_map"]
reduce_partial_map = meta["reduce_partial_map"]
sm_scale = float(config.get("sm_scale", SM_SCALE))
pages_per_seq = (kv_seq_len + page_size - 1) // page_size
output_buffers = [
torch.empty((num_tokens, nhead, v_head_dim), dtype=torch.bfloat16, device=device)
for _ in range(_OUTPUT_CACHE_SLOTS)
]
output_index = 0
num_output_buffers = len(output_buffers)
if Q_DTYPE == "fp8":
q_cache_slots = _EXACT_Q_CACHE_SLOTS_BY_SHAPE.get(shape, _EXACT_SAFE_Q_CACHE_SLOTS)
q_buffers = [
torch.empty((num_tokens, nhead, qk_head_dim), dtype=FP8_DTYPE, device=device)
for _ in range(q_cache_slots)
]
q_index = 0
num_q_buffers = len(q_buffers)
unit_q_scale = _get_unit_q_scale(device)
else:
q_buffers = []
q_index = 0
num_q_buffers = 0
unit_q_scale = None
if kv_seq_len % page_size == 0:
kv_view_shape = (batch_size * pages_per_seq, page_size, nhead_kv, qk_head_dim)
def reshape_kv(kv_buffer: torch.Tensor) -> torch.Tensor:
return kv_buffer.view(*kv_view_shape)
else:
def reshape_kv(kv_buffer: torch.Tensor) -> torch.Tensor:
return _reshape_paged_kv(kv_buffer, batch_size, kv_seq_len, page_size, nhead_kv, qk_head_dim)
def runner(q: torch.Tensor, kv_buffer: torch.Tensor, kv_scale: torch.Tensor | None) -> torch.Tensor:
nonlocal output_index, q_index
if Q_DTYPE == "fp8":
q_input = q_buffers[q_index]
q_input.copy_(q)
q_index += 1
if q_index == num_q_buffers:
q_index = 0
q_scale = unit_q_scale
else:
q_input = q
q_scale = None
out = output_buffers[output_index]
output_index += 1
if output_index == num_output_buffers:
output_index = 0
mla_decode_fwd(
q_input.reshape(-1, nhead, qk_head_dim),
reshape_kv(kv_buffer),
out,
qo_indptr,
kv_page_indptr,
kv_page_indices,
kv_last_page_len,
q_seq_len,
page_size=page_size,
nhead_kv=nhead_kv,
sm_scale=sm_scale,
logit_cap=0.0,
num_kv_splits=num_kv_splits,
q_scale=q_scale,
kv_scale=kv_scale,
intra_batch_mode=intra_batch_mode,
work_meta_data=work_meta_data,
work_indptr=work_indptr,
work_info_set=work_info_set,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
)
return out
return runner
def _get_exact_shape_runner(shape: tuple[int, int], kv_dtype: torch.dtype, device: torch.device):
bank_key = (device.index or 0, kv_dtype)
bank = _EXACT_SHAPE_RUNNER_BANK.get(bank_key)
if bank is None:
bank = {}
_EXACT_SHAPE_RUNNER_BANK[bank_key] = bank
runner = bank.get(shape)
if runner is None:
config = _EXACT_SHAPE_CONFIGS.get(shape)
if config is None:
return None
runner = _make_exact_shape_runner(shape, kv_dtype, device)
bank[shape] = runner
return runner
def custom_kernel(data: input_t) -> output_t:
q, kv_data, _qo_indptr, _kv_indptr, config = data
if int(config["q_seq_len"]) != 1:
raise ValueError(f"Expected decode q_seq_len=1, got {config['q_seq_len']}")
shape = (int(config["batch_size"]), int(config["kv_seq_len"]))
if KV_DTYPE == "fp8":
kv_input, kv_scale = kv_data["fp8"]
else:
kv_input, kv_scale = kv_data["bf16"], None
if shape in _FAST_PATH_SHAPES:
runner = _get_exact_shape_runner(shape, kv_input.dtype, q.device)
if runner is not None:
return runner(q, kv_input, kv_scale)
return _get_persistent_runner(config, kv_input.dtype, q.device)(q, kv_input, kv_scale)
scrolls · 589 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON