submission 656193
谢书骁 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 147 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-656193?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:c97ee437d05de43c7743506ce6e8e515611d3266f9d819308fd089eebc0f13d0
license declaredunknown
license concludedunknown
authors谢书骁
imported2026-08-26
Kernel source
submission.py147 lines
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
try:
from aiter import dynamic_per_tensor_quant
_HAS_AITER_QUANT = True
except ImportError:
_HAS_AITER_QUANT = False
NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
NUM_KV_SPLITS = 32
FP8_DTYPE = aiter_dtypes.fp8
FP8_MAX = torch.finfo(FP8_DTYPE).max
FP8_MIN = torch.finfo(FP8_DTYPE).min
# ---------------------------------------------------------------------------
# Flat cache: avoid dict lookups in hot path
# ---------------------------------------------------------------------------
class _CachedState:
__slots__ = [
'kv_indices', 'kv_last_page_len', 'max_q_len', 'o',
'q_fp8_buf', 'q_scale_buf', 'q_cache_ptr',
'kv_cache_ptr', 'kv_buffer_4d',
'total_kv',
'work_meta_data', 'work_indptr', 'work_info_set',
'reduce_indptr', 'reduce_final_map', 'reduce_partial_map',
]
_shape_cache = {}
def _build_shape_cache(batch_size, q_seq_len, kv_seq_len, total_q, total_kv, qo_indptr, kv_indptr):
nq = NUM_HEADS
nkv = NUM_KV_HEADS
s = _CachedState()
s.total_kv = total_kv
s.kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
s.kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
s.max_q_len = q_seq_len
s.q_cache_ptr = None
s.kv_cache_ptr = None
s.kv_buffer_4d = None
info = get_mla_metadata_info_v1(
batch_size, s.max_q_len, nq, FP8_DTYPE, FP8_DTYPE,
is_sparse=False, fast_mode=False,
num_kv_splits=NUM_KV_SPLITS, intra_batch_mode=True,
)
work = [torch.empty(sz, dtype=dt, device="cuda") for sz, dt in info]
s.work_meta_data, s.work_indptr, s.work_info_set, \
s.reduce_indptr, s.reduce_final_map, s.reduce_partial_map = work
get_mla_metadata_v1(
qo_indptr, kv_indptr, s.kv_last_page_len,
nq // nkv, nkv, True,
s.work_meta_data, s.work_info_set, s.work_indptr,
s.reduce_indptr, s.reduce_final_map, s.reduce_partial_map,
page_size=PAGE_SIZE,
kv_granularity=max(PAGE_SIZE, 16),
max_seqlen_qo=s.max_q_len,
uni_seqlen_qo=s.max_q_len,
fast_mode=False,
max_split_per_batch=NUM_KV_SPLITS,
intra_batch_mode=True,
dtype_q=FP8_DTYPE,
dtype_kv=FP8_DTYPE,
)
s.o = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
s.q_fp8_buf = torch.empty((total_q, nq, QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda")
s.q_scale_buf = torch.empty(1, dtype=torch.float32, device="cuda")
return s
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
# Shape cache lookup — use batch_size + kv_seq_len as key
bs = config["batch_size"]
kvlen = config["kv_seq_len"]
shape_key = (bs, kvlen)
s = _shape_cache.get(shape_key)
if s is None:
total_q = q.shape[0]
total_kv = int(kv_indptr[-1].item())
s = _build_shape_cache(bs, config["q_seq_len"], kvlen, total_q, total_kv, qo_indptr, kv_indptr)
_shape_cache[shape_key] = s
# Q quantization: skip if same data
q_ptr = q.data_ptr()
if q_ptr != s.q_cache_ptr:
if _HAS_AITER_QUANT:
dynamic_per_tensor_quant(s.q_fp8_buf, q, s.q_scale_buf)
else:
amax = q.abs().amax().clamp(min=1e-12)
sc = amax / FP8_MAX
s.q_fp8_buf.copy_((q / sc).clamp(min=FP8_MIN, max=FP8_MAX).to(FP8_DTYPE))
s.q_scale_buf.fill_(sc.item())
s.q_cache_ptr = q_ptr
# KV view: cache if same buffer
kv_fp8_tuple = kv_data["fp8"]
kv_buffer_fp8 = kv_fp8_tuple[0]
kv_ptr = kv_buffer_fp8.data_ptr()
if kv_ptr != s.kv_cache_ptr:
s.kv_buffer_4d = kv_buffer_fp8.view(s.total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)
s.kv_cache_ptr = kv_ptr
mla_decode_fwd(
s.q_fp8_buf,
s.kv_buffer_4d,
s.o,
qo_indptr,
kv_indptr,
s.kv_indices,
s.kv_last_page_len,
s.max_q_len,
page_size=PAGE_SIZE,
nhead_kv=NUM_KV_HEADS,
sm_scale=SM_SCALE,
logit_cap=0.0,
num_kv_splits=NUM_KV_SPLITS,
q_scale=s.q_scale_buf,
kv_scale=kv_fp8_tuple[1],
intra_batch_mode=True,
work_meta_data=s.work_meta_data,
work_indptr=s.work_indptr,
work_info_set=s.work_info_set,
reduce_indptr=s.reduce_indptr,
reduce_final_map=s.reduce_final_map,
reduce_partial_map=s.reduce_partial_map,
)
return s.o
scrolls · 147 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