submission 653127
mocimex265 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 329 lines, June 9 Researcher Reciprocity License v1.0.
test_v96_leaderboard_safe.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-653127?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:40aacdaa82e74a5b808caa9422ab228350d4f22282d40c33ce13300804ea75c2
license declaredunknown
license concludedunknown
authorsmocimex265
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
Target: s4/s6 — revert to persistent (nonpers fails leaderboard for kv=8192)Kernel source
test_v96_leaderboard_safe.py329 lines
"""
test_v96_leaderboard_safe: Revert nonpers to kv<=1024 only (fix leaderboard s4 failure)
Base: test_v94_gran64_all.py
Direction: NEW — fix leaderboard correctness
Target: s4/s6 — revert to persistent (nonpers fails leaderboard for kv=8192)
Change: Restore kv_seq_len<=1024 condition for nonpers fp8 tier. s4/s6 back to
persistent with kv_gran=64. Keep nonpers for s3/s5 (kv<=1024, proven safe).
v94 leaderboard failed on s4 (batch=32, kv=8192) — nonpers + custom reduce
breaks for kv=8192 in leaderboard mode (same pattern as v77-v78).
Rationale: v94 leaderboard failed on s4. Must revert kv>1024 to persistent.
Scale: INCREMENTAL
"""
import torch
import aiter
import triton
import triton.language as tl
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
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 = {}
# ========================================================================
# Custom Triton reduce kernel (replaces aiter's _fwd_kernel_stage2_asm)
# Simple, stateless, no caching — should be leaderboard-safe.
# ========================================================================
@triton.jit
def _custom_reduce_kernel(
logits_ptr, # float32 (total_q, num_splits, num_heads, v_dim)
lse_ptr, # float32 (total_q, num_splits, num_heads, 1)
output_ptr, # bf16 (total_q, num_heads, v_dim)
total_q: tl.int32,
num_splits: tl.constexpr,
num_heads: tl.constexpr,
v_dim: tl.constexpr,
BLOCK_V: tl.constexpr,
):
# Grid: (total_q, num_heads, cdiv(v_dim, BLOCK_V))
q_idx = tl.program_id(0)
head_idx = tl.program_id(1)
v_block = tl.program_id(2)
v_offs = v_block * BLOCK_V + tl.arange(0, BLOCK_V)
v_mask = v_offs < v_dim
# Strides for logits: (total_q, num_splits, num_heads, v_dim)
logits_q_stride = num_splits * num_heads * v_dim
logits_s_stride = num_heads * v_dim
logits_h_stride = v_dim
# Strides for lse: (total_q, num_splits, num_heads, 1)
lse_q_stride = num_splits * num_heads
lse_s_stride = num_heads
# Find max LSE across splits for numerical stability
max_lse = tl.full((), float("-inf"), dtype=tl.float32)
for s in range(num_splits):
lse_val = tl.load(lse_ptr + q_idx * lse_q_stride + s * lse_s_stride + head_idx)
max_lse = tl.maximum(max_lse, lse_val)
# Weighted sum of logits across splits
acc = tl.zeros((BLOCK_V,), dtype=tl.float32)
weight_sum = tl.full((), 0.0, dtype=tl.float32)
for s in range(num_splits):
lse_val = tl.load(lse_ptr + q_idx * lse_q_stride + s * lse_s_stride + head_idx)
w = tl.exp(lse_val - max_lse)
weight_sum += w
logits_base = logits_ptr + q_idx * logits_q_stride + s * logits_s_stride + head_idx * logits_h_stride
vals = tl.load(logits_base + v_offs, mask=v_mask, other=0.0)
acc += vals * w
# Normalize
acc = acc / weight_sum
# Store as bf16
out_base = output_ptr + q_idx * num_heads * v_dim + head_idx * v_dim
tl.store(out_base + v_offs, acc.to(tl.bfloat16), mask=v_mask)
# ========================================================================
# Cache functions
# ========================================================================
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_nonpers_fp8(batch_size, kv_seq_len, total_q):
"""Non-persistent fp8 with float32 logits and custom reduce."""
key = ("npfp8_cr", 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, FP8_DTYPE)
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")
# Always float32 logits — never alias to output
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,
"q_fp8": q_fp8, "q_scale": q_scale,
}
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]
# ========================================================================
# Main dispatch
# ========================================================================
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,
)
# Use custom reduce for bf16 nonpers too (for consistency)
num_splits = c["num_kv_splits"]
BLOCK_V = 128
n_v_blocks = triton.cdiv(V_HEAD_DIM, BLOCK_V)
_custom_reduce_kernel[(total_q, NUM_HEADS, n_v_blocks)](
c["logits"], c["attn_lse"], c["o"],
total_q,
num_splits=num_splits,
num_heads=NUM_HEADS,
v_dim=V_HEAD_DIM,
BLOCK_V=BLOCK_V,
num_warps=4,
)
return c["o"]
# ---- Tier 2: nonpers fp8 with custom reduce (s3, s5 ONLY) ----
# Only kv<=1024 is safe for nonpers in leaderboard mode.
# kv=8192 nonpers FAILS leaderboard (v94 failed on s4).
elif kv_seq_len <= 1024 and batch_size <= 64:
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)
c = _ensure_cache_nonpers_fp8(batch_size, kv_seq_len, total_q)
# Fast FP8 quant: copy_ cast (scale=1.0)
q_2d = q.view(total_q, NUM_HEADS * QK_HEAD_DIM)
c["q_fp8"].copy_(q_2d)
# Stage 1: non-persistent fp8 attention
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"],
c["num_kv_splits_indptr"],
None, None, None,
1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
c["logits"], c["attn_lse"], c["o"],
c["q_scale"], kv_scale,
)
# Stage 2: CUSTOM reduce (replaces aiter's _fwd_kernel_stage2_asm)
num_splits = c["num_kv_splits"]
BLOCK_V = 128
n_v_blocks = triton.cdiv(V_HEAD_DIM, BLOCK_V)
_custom_reduce_kernel[(total_q, NUM_HEADS, n_v_blocks)](
c["logits"], c["attn_lse"], c["o"],
total_q,
num_splits=num_splits,
num_heads=NUM_HEADS,
v_dim=V_HEAD_DIM,
BLOCK_V=BLOCK_V,
num_warps=4,
)
return c["o"]
# ---- Tier 3: fp8 persistent (s4, s6, s7, s8) ----
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:
splits, fast_mode = 4, False # s8
elif total_kv >= 300000:
splits, fast_mode = 8, False # s6
elif batch_size >= 256:
splits, fast_mode = 1, False # s7
elif batch_size >= 64:
splits, fast_mode = 4, False # s5 won't reach here (kv<=1024 → tier 2)
else:
splits, fast_mode = 32, True # s4
kv_gran = 64 # use 64 for all persistent shapes (was 16 for kv<=1024)
c = _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, splits, fast_mode, kv_gran)
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 · 329 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