submission 629089
michael ma · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 91 lines, June 9 Researcher Reciprocity License v1.0.
submission4.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-629089?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:921eef02b58e28f8b1747f4f1e952911560f094178c0a36671d03b7713796187
license declaredunknown
license concludedunknown
authorsmichael ma
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
基于理论模型选择最优 split-K 数量:Kernel source
submission4.py91 lines
import torch
from aiter.mla import mla_decode_fwd
from reference import (
NUM_HEADS, NUM_KV_HEADS, QK_HEAD_DIM, V_HEAD_DIM,
SM_SCALE, PAGE_SIZE, _make_mla_decode_metadata, quantize_fp8
)
def select_num_kv_splits(kv_seq_len: int, batch_size: int) -> int:
"""
基于理论模型选择最优 split-K 数量:
S_opt ≈ sqrt(L * B_kv / overhead)
其中 B_kv = 576 (FP8 每 token 字节数),overhead ≈ 2KB(合并开销)
结果再根据 batch 规模调整,并取 2 的幂以便 kernel 调度。
"""
L = kv_seq_len
B_kv = 576
overhead = 2048 # 2KB
s_theory = int((L * B_kv / overhead) ** 0.5)
# 限制范围
s = max(1, min(128, s_theory))
# 针对大 batch 长序列可适当增加
if batch_size >= 64 and L >= 8192:
s = max(s, 64)
# 量化到 2 的幂
if s < 16:
return 16
elif s < 32:
return 32
elif s < 64:
return 64
else:
return 128
def custom_kernel(data):
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
kv_seq_len = config["kv_seq_len"]
# 自适应 split-K
num_kv_splits = select_num_kv_splits(kv_seq_len, batch_size)
# 使用 FP8 加速(aiter 最优路径)
kv_buffer, kv_scale = kv_data["fp8"]
q_fp8, q_scale = quantize_fp8(q)
# 准备 4D KV buffer
kv_buffer_4d = kv_buffer.view(kv_buffer.shape[0], PAGE_SIZE, NUM_KV_HEADS, kv_buffer.shape[-1])
total_kv_len = int(kv_indptr[-1].item())
kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
# 构建 metadata(不使用 fast_mode,保持与参考一致)
meta = _make_mla_decode_metadata(
batch_size,
config["q_seq_len"],
NUM_HEADS,
NUM_KV_HEADS,
q_fp8.dtype,
kv_buffer.dtype,
qo_indptr,
kv_indptr,
kv_last_page_len,
num_kv_splits=num_kv_splits,
)
o = torch.empty(
(q.shape[0], NUM_HEADS, V_HEAD_DIM),
dtype=torch.bfloat16, device="cuda"
)
mla_decode_fwd(
q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM),
kv_buffer_4d,
o,
qo_indptr,
kv_indptr,
kv_indices,
kv_last_page_len,
config["q_seq_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=q_scale,
kv_scale=kv_scale,
intra_batch_mode=True,
**meta,
)
return oscrolls · 91 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