submission 738884
Behzod12312121 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 302 lines, June 9 Researcher Reciprocity License v1.0.
submission_quant_hip_splitmap_s3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-738884?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:4e193b564639ef7894c730f9beb51ea542dc799d744ef20958f60d80c383b768
license declaredunknown
license concludedunknown
authorsBehzod12312121
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
surface. The persistent decode path and the best measured split schedule stayKernel source
submission_quant_hip_splitmap_s3.py302 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
Mixed-MLA contingency candidate: use AITER's per_tensor_quant_hip helper.
This tests whether the library's higher-level quant helper performs differently
enough from the direct compiled-op wrapper to matter on the official benchmark
surface. The persistent decode path and the best measured split schedule stay
the same as the current MLA best file.
"""
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
from aiter.ops.quant import per_tensor_quant_hip
NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM
V_HEAD_DIM = KV_LORA_RANK
SM_SCALE = 1.0 / (QK_HEAD_DIM**0.5)
PAGE_SIZE = 1
NUM_KV_SPLITS = 32
FP8_DTYPE = aiter_dtypes.fp8
SPLIT_OVERRIDES = {
(1, 512): 8,
(1, 2048): 12,
(1, 8192): 20,
(8, 512): 8,
(8, 2048): 16,
(8, 8192): 24,
(16, 512): 12,
(16, 2048): 20,
(16, 8192): 28,
(32, 512): 16,
(32, 2048): 24,
(32, 8192): 32,
(64, 512): 20,
(64, 2048): 28,
(64, 8192): 36,
}
KV_GRANULARITY = 32
_WORKSPACE_CACHE: dict[tuple[object, ...], tuple[torch.Tensor, ...]] = {}
_INDPTR_CACHE: dict[tuple[object, ...], tuple[torch.Tensor, torch.Tensor, torch.Tensor]] = {}
_KV_INDEX_CACHE: dict[tuple[object, ...], torch.Tensor] = {}
_METADATA_CACHE: dict[tuple[object, ...], dict[str, torch.Tensor]] = {}
_OUTPUT_CACHE: dict[tuple[object, ...], torch.Tensor] = {}
def quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
if not tensor.is_contiguous():
tensor = tensor.contiguous()
return per_tensor_quant_hip(tensor, quant_dtype=FP8_DTYPE)
def _device_cache_key(device: torch.device) -> tuple[str, int]:
return device.type, -1 if device.index is None else device.index
def _get_metadata_workspace(
batch_size: int,
max_q_len: int,
nhead: int,
nhead_kv: int,
q_dtype: torch.dtype,
kv_dtype: torch.dtype,
num_kv_splits: int,
device: torch.device,
) -> tuple[torch.Tensor, ...]:
cache_key = (
_device_cache_key(device),
batch_size,
max_q_len,
nhead,
nhead_kv,
q_dtype,
kv_dtype,
num_kv_splits,
)
cached = _WORKSPACE_CACHE.get(cache_key)
if cached is None:
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=True,
)
cached = tuple(torch.empty(shape, dtype=dtype, device=device) for shape, dtype in info)
_WORKSPACE_CACHE[cache_key] = cached
return cached
def _get_kv_indices(total_kv_len: int, device: torch.device) -> torch.Tensor:
cache_key = (_device_cache_key(device), total_kv_len)
cached = _KV_INDEX_CACHE.get(cache_key)
if cached is None:
cached = torch.arange(total_kv_len, dtype=torch.int32, device=device)
_KV_INDEX_CACHE[cache_key] = cached
return cached
def _get_uniform_indptrs(
batch_size: int,
q_seq_len: int,
kv_seq_len: int,
device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
cache_key = (_device_cache_key(device), batch_size, q_seq_len, kv_seq_len)
cached = _INDPTR_CACHE.get(cache_key)
if cached is None:
steps = torch.arange(batch_size + 1, dtype=torch.int32, device=device)
cached = (
steps * q_seq_len,
steps * kv_seq_len,
torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device),
)
_INDPTR_CACHE[cache_key] = cached
return cached
def _get_output_buffer(shape: tuple[int, int, int], device: torch.device) -> torch.Tensor:
cache_key = (_device_cache_key(device), shape)
cached = _OUTPUT_CACHE.get(cache_key)
if cached is None:
cached = torch.empty(shape, dtype=torch.bfloat16, device=device)
_OUTPUT_CACHE[cache_key] = cached
return cached
def _select_num_kv_splits(batch_size: int, kv_seq_len: int) -> int:
override = SPLIT_OVERRIDES.get((batch_size, kv_seq_len))
if override is not None:
return override
total_kv = batch_size * kv_seq_len
if total_kv >= 1_000_000:
return 32
if total_kv >= 131_072:
return 24
return 16
def make_mla_decode_metadata(
batch_size: int,
max_q_len: int,
kv_seq_len: int,
nhead: int,
nhead_kv: int,
q_dtype: torch.dtype,
kv_dtype: torch.dtype,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
kv_last_page_len: torch.Tensor,
num_kv_splits: int = NUM_KV_SPLITS,
):
cache_key = (
_device_cache_key(qo_indptr.device),
batch_size,
max_q_len,
kv_seq_len,
nhead,
nhead_kv,
q_dtype,
kv_dtype,
num_kv_splits,
)
cached = _METADATA_CACHE.get(cache_key)
if cached is not None:
return cached
(
work_metadata,
work_indptr,
work_info_set,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
) = _get_metadata_workspace(
batch_size,
max_q_len,
nhead,
nhead_kv,
q_dtype,
kv_dtype,
num_kv_splits,
qo_indptr.device,
)
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=max(PAGE_SIZE, 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=True,
dtype_q=q_dtype,
dtype_kv=kv_dtype,
)
cached = {
"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,
}
_METADATA_CACHE[cache_key] = cached
return cached
def custom_kernel(data: input_t) -> output_t:
q, kv_data, _, _, config = data
batch_size = config["batch_size"]
nq = config["num_heads"]
nkv = config["num_kv_heads"]
dq = config["qk_head_dim"]
dv = config["v_head_dim"]
q_seq_len = config["q_seq_len"]
kv_seq_len = config["kv_seq_len"]
qo_indptr, kv_indptr, kv_last_page_len = _get_uniform_indptrs(
batch_size,
q_seq_len,
kv_seq_len,
q.device,
)
q_fp8, q_scale = quantize_fp8(q)
kv_buffer_fp8, kv_scale = kv_data["fp8"]
total_kv_len = kv_buffer_fp8.shape[0]
kv_indices = _get_kv_indices(total_kv_len, q.device)
kv_buffer_4d = kv_buffer_fp8.reshape(total_kv_len, PAGE_SIZE, nkv, kv_buffer_fp8.shape[-1])
num_kv_splits = _select_num_kv_splits(batch_size, kv_seq_len)
meta = make_mla_decode_metadata(
batch_size,
q_seq_len,
kv_seq_len,
nq,
nkv,
q_fp8.dtype,
kv_buffer_4d.dtype,
qo_indptr,
kv_indptr,
kv_last_page_len,
num_kv_splits=num_kv_splits,
)
o = _get_output_buffer((q.shape[0], nq, dv), q.device)
mla_decode_fwd(
q_fp8.reshape(-1, nq, dq),
kv_buffer_4d,
o,
qo_indptr,
kv_indptr,
kv_indices,
kv_last_page_len,
q_seq_len,
page_size=PAGE_SIZE,
nhead_kv=nkv,
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=True,
**meta,
)
return o
scrolls · 302 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