submission 718159
tangzhanshuo · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 207 lines, June 9 Researcher Reciprocity License v1.0.
mla_ref.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-718159?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:ff3d79514137e100e1083ba5ff7b99f4e0c75bc00d567179972dfb31a07ebc8e
license declaredunknown
license concludedunknown
authorstangzhanshuo
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
num_warps=4,persistent-kernel
def _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, persistent_splits, fast_mode, kv_gran=16):stages = 2
num_stages=2,Kernel source
mla_ref.py207 lines
"""
test_v143_s6_splits4: s6 splits 8→4 (continue reduce optimization pattern)
Base: test.py (v142)
Direction: NEW — s6 splits tuning
Target: s6 (64,8192) — reduce overhead with kv=8192
Change: s6 splits 8→4. batch=64 × splits=4 = 256 programs (100% CU fill).
Follows s5 splits reduction pattern (v142 +8.5%).
Rationale: v140 profile s8 reduce=3.3us. s6 with splits=8 has more reduce overhead.
Scale: INCREMENTAL
"""
import torch
import aiter
import triton
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.mla import get_meta_param, _fwd_kernel_stage2_asm
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
FP8_DTYPE = aiter_dtypes.fp8
_cache = {}
def _ensure_cache_nonpers_bf16(batch_size, kv_seq_len, total_q):
key = ("npbf16", batch_size, kv_seq_len)
if key in _cache:
return _cache[key]
nq = NUM_HEADS
total_kv = batch_size * kv_seq_len
kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")
kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
num_kv_splits, num_kv_splits_indptr = get_meta_param(None, batch_size, total_kv, nq, 1, torch.bfloat16)
o = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
logits = torch.empty((total_q, num_kv_splits, nq, V_HEAD_DIM), dtype=torch.float32, device="cuda")
attn_lse = torch.empty((total_q, num_kv_splits, nq, 1), dtype=torch.float32, device="cuda")
_cache[key] = {
"kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,
"num_kv_splits": num_kv_splits, "num_kv_splits_indptr": num_kv_splits_indptr,
"logits": logits, "attn_lse": attn_lse, "o": o,
}
return _cache[key]
def _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, persistent_splits, fast_mode, kv_gran=16):
key = ("pfp8", batch_size, kv_seq_len, persistent_splits, fast_mode, kv_gran)
if key in _cache:
return _cache[key]
max_q_len = 1
nq, nkv = NUM_HEADS, NUM_KV_HEADS
total_kv = batch_size * kv_seq_len
kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")
kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
info = get_mla_metadata_info_v1(
batch_size, max_q_len, nq, FP8_DTYPE, FP8_DTYPE,
is_sparse=False, fast_mode=fast_mode,
num_kv_splits=persistent_splits, intra_batch_mode=True,
)
work = [torch.empty(s, dtype=t, device="cuda") for s, t 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,
nq // nkv, nkv, 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_gran),
max_seqlen_qo=max_q_len,
uni_seqlen_qo=max_q_len,
fast_mode=fast_mode,
max_split_per_batch=persistent_splits,
intra_batch_mode=True,
dtype_q=FP8_DTYPE,
dtype_kv=FP8_DTYPE,
)
num_partials = reduce_partial_map.size(0)
logits = torch.empty((num_partials, 1, nq, V_HEAD_DIM), dtype=torch.float32, device="cuda")
attn_lse = torch.empty((num_partials, 1, nq, 1), dtype=torch.float32, device="cuda")
o = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
q_fp8 = torch.empty((total_q, nq * QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda")
q_scale = torch.ones(1, dtype=torch.float32, device="cuda")
_cache[key] = {
"kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,
"work_metadata": 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,
"logits": logits, "attn_lse": attn_lse, "o": o,
"q_fp8": q_fp8, "q_scale": q_scale,
"num_partials": num_partials,
}
return _cache[key]
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
kv_seq_len = config["kv_seq_len"]
total_q = q.shape[0]
total_kv = batch_size * kv_seq_len
# ---- Tier 1: batch<=4 -> bf16/bf16 non-persistent (s1, s2) ----
if batch_size <= 4:
kv_bf16 = kv_data["bf16"]
kv_4d = kv_bf16.view(total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)
c = _ensure_cache_nonpers_bf16(batch_size, kv_seq_len, total_q)
aiter.mla_decode_stage1_asm_fwd(
q.view(-1, NUM_HEADS, QK_HEAD_DIM), kv_4d,
qo_indptr, kv_indptr, c["kv_indices"], c["kv_last_page_len"],
c["num_kv_splits_indptr"],
None, None, None,
1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
c["logits"], c["attn_lse"], c["o"],
None, None,
)
Lv = V_HEAD_DIM
BLOCK_DV = triton.next_power_of_2(Lv)
_fwd_kernel_stage2_asm[(batch_size, NUM_HEADS)](
c["logits"], c["attn_lse"], c["o"],
qo_indptr, kv_indptr, c["num_kv_splits_indptr"],
c["attn_lse"].stride(0), c["attn_lse"].stride(2), c["attn_lse"].stride(1),
c["o"].stride(0), c["o"].stride(1),
MAYBE_FINAL_OUT=True,
BATCH_NUM=batch_size,
BLOCK_DV=BLOCK_DV,
Lv=Lv,
mgc=64,
num_warps=4,
num_stages=2,
waves_per_eu=4,
)
return c["o"]
# ---- Tier 2: ALL fp8 shapes -> persistent (s3-s8) ----
# Non-persistent fp8 was faster but fails leaderboard correctness (v77, v78).
# Persistent + mla_reduce_v1 is the only leaderboard-safe fp8 path.
else:
kv_buffer_fp8, kv_scale = kv_data["fp8"]
kv_buffer_4d = kv_buffer_fp8.view(total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)
# Per-shape split tuning
if total_kv >= 1000000:
# s8 (256, 8192) -- splits=4 (from v77)
splits, fast_mode = 4, False
elif total_kv >= 300000:
# s6 (64, 8192) -- splits=4 (from 8, 64*4=256 programs = 100% CU fill)
splits, fast_mode = 4, False
elif batch_size >= 256:
# s7 (256, 1024) -- splits=4 with kv_gran=64 (v131 LB-safe config)
# splits=1+kv_gran=64 FAILED LB in v136. splits=4 gives 1024 programs.
splits, fast_mode = 4, False
elif batch_size >= 64:
# s5 (64, 1024) -- splits=2 (from 4, reduce=7us → ~3.5us, 128 programs = 50% CU)
splits, fast_mode = 2, False
else:
# s3 (32, 1024) and s4 (32, 8192)
if kv_seq_len <= 1024:
splits, fast_mode = 4, True # s3: reduced from 8 to 4
else:
splits, fast_mode = 32, True # s4
# Use kv_granularity=64 for ALL persistent shapes (matches v131 LB-safe config)
# v131 (kv_gran=64 all) PASSED LB at 58.0us. v137/v138 (kv_gran=16 for s3/s5)
# FAILED LB on s3. kv_gran=64 is required for LB correctness.
kv_gran = 64
c = _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, splits, fast_mode, kv_gran)
# Fast FP8 quant: copy_ cast (scale=1.0) -- from v63
q_2d = q.view(total_q, NUM_HEADS * QK_HEAD_DIM)
c["q_fp8"].copy_(q_2d)
aiter.mla_decode_stage1_asm_fwd(
c["q_fp8"].view(-1, NUM_HEADS, QK_HEAD_DIM), kv_buffer_4d,
qo_indptr, kv_indptr, c["kv_indices"], c["kv_last_page_len"],
None, c["work_metadata"], c["work_indptr"], c["work_info_set"],
1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
c["logits"], c["attn_lse"], c["o"],
c["q_scale"], kv_scale,
)
aiter.mla_reduce_v1(
c["logits"], c["attn_lse"],
c["reduce_indptr"], c["reduce_final_map"], c["reduce_partial_map"],
1, c["o"], None,
)
return c["o"]
scrolls · 207 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