submission 736894
NinoHeather · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 321 lines, June 9 Researcher Reciprocity License v1.0.
my_submission_refact.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-736894?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:a67cc2fd250671ae326070458a9d2a5c5713d73e44c1f1ec304e8bd1b63a4f39
license declaredunknown
license concludedunknown
authorsNinoHeather
imported2026-08-15
Kernel source
my_submission_refact.py321 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
import os
import torch
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 import mla_decode_stage1_asm_fwd, mla_reduce_v1
from aiter.mla import mla_decode_fwd
HAS_STATIC_QUANT = False
try:
# Prefer JIT quant op when available.
from aiter.jit.module_quant import static_per_tensor_quant
HAS_STATIC_QUANT = True
except Exception:
try:
from aiter.ops.quant import static_per_tensor_quant
HAS_STATIC_QUANT = True
except Exception:
pass
try:
from aiter.ops.quant import dynamic_per_tensor_quant
HAS_DYNAMIC_QUANT = True
except ImportError:
HAS_DYNAMIC_QUANT = False
FP8_T = aiter_dtypes.fp8
BF16_T = torch.bfloat16
N_HEADS = 16
N_KV_HEADS = 1
QK_DIM = 576
V_DIM = 512
ATTN_SCALE = 1.0 / (QK_DIM ** 0.5)
STATIC_Q_SCALE = torch.tensor([0.1], dtype=torch.float32, device="cuda")
# (batch_size, kv_seq_len) -> (page_size, num_kv_splits, is_sparse, kv_granularity_override)
SHAPE_POLICY = {
(4, 1024): (2, 8, False, None),
(4, 8192): (8, 16, False, None),
(32, 1024): (2, 8, False, None),
(32, 8192): (8, 12, False, None),
(64, 1024): (2, 1, False, 4),
(64, 8192): (8, 12, False, None),
(256, 1024): (2, 1, True, 2),
(256, 8192): (8, 8, False, None),
}
NP_8K_PAGE = {
4: 2048,
32: 1024,
64: 512,
256: 2048,
}
FWD_A8W8_1K_SHAPES = {
(4, 1024),
(32, 1024),
}
def _build_shape_state(
bs: int,
kv_len: int,
page_size: int,
splits: int,
sparse: bool,
kv_granularity_override: int | None,
) -> dict:
pages_per_req = kv_len // page_size
total_pages = bs * pages_per_req
qo_prefix = torch.arange(bs + 1, dtype=torch.int32, device="cuda")
kv_pages_prefix = torch.arange(bs + 1, dtype=torch.int32, device="cuda") * pages_per_req
kv_last_page_len = torch.full((bs,), page_size, dtype=torch.int32, device="cuda")
kv_indices = torch.arange(total_pages, dtype=torch.int32, device="cuda")
meta_info = get_mla_metadata_info_v1(
bs,
1,
N_HEADS,
FP8_T,
FP8_T,
is_sparse=sparse,
fast_mode=True,
num_kv_splits=splits,
intra_batch_mode=True,
)
work_meta = [torch.empty(sz, dtype=dt, device="cuda") for sz, dt in meta_info]
wm, wi, wis, ri, rfm, rpm = work_meta
kv_granularity = kv_granularity_override if kv_granularity_override is not None else max(page_size, 16)
get_mla_metadata_v1(
qo_prefix,
kv_pages_prefix,
kv_last_page_len,
N_HEADS // N_KV_HEADS,
N_KV_HEADS,
False,
wm,
wis,
wi,
ri,
rfm,
rpm,
page_size=page_size,
kv_granularity=kv_granularity,
max_seqlen_qo=1,
uni_seqlen_qo=1,
fast_mode=True,
max_split_per_batch=splits,
intra_batch_mode=True,
dtype_q=FP8_T,
dtype_kv=FP8_T,
)
partial_rows = rpm.size(0)
return {
"page_size": page_size,
"splits": splits,
"qo_prefix": qo_prefix,
"kv_pages_prefix": kv_pages_prefix,
"kv_last_page_len": kv_last_page_len,
"kv_indices": kv_indices,
"wm": wm,
"wi": wi,
"wis": wis,
"ri": ri,
"rfm": rfm,
"rpm": rpm,
"partial_logits": torch.empty((partial_rows, 1, N_HEADS, V_DIM), dtype=torch.float32, device="cuda"),
"partial_lse": torch.empty((partial_rows, 1, N_HEADS, 1), dtype=torch.float32, device="cuda"),
"out_bf16": torch.empty((bs, N_HEADS, V_DIM), dtype=BF16_T, device="cuda"),
"q_fp8": torch.empty((bs, N_HEADS, QK_DIM), dtype=FP8_T, device="cuda"),
"q_scale": STATIC_Q_SCALE.clone(),
}
def _build_fwd_a8w8_1k_state(bs: int, kv_len: int, page_size: int = 2) -> dict:
pages_per_req = kv_len // page_size
total_pages = bs * pages_per_req
return {
"page_size": page_size,
"kv_indices": torch.arange(total_pages, dtype=torch.int32, device="cuda"),
"kv_pages_prefix": torch.arange(bs + 1, dtype=torch.int32, device="cuda") * pages_per_req,
"kv_last_page_lens": torch.full((bs,), page_size, dtype=torch.int32, device="cuda"),
"out_bf16": torch.empty((bs, N_HEADS, V_DIM), dtype=BF16_T, device="cuda"),
"q_fp8": torch.empty((bs, N_HEADS, QK_DIM), dtype=FP8_T, device="cuda"),
"q_scale": STATIC_Q_SCALE.clone(),
}
def _build_np8k_state(bs: int, kv_len: int, page_size: int) -> dict:
pages_per_req = kv_len // page_size
total_pages = bs * pages_per_req
return {
"page_size": page_size,
"kv_indices": torch.arange(total_pages, dtype=torch.int32, device="cuda"),
"kv_pages_prefix": torch.arange(bs + 1, dtype=torch.int32, device="cuda") * pages_per_req,
"kv_last_page_lens": torch.full((bs,), page_size, dtype=torch.int32, device="cuda"),
"num_kv_splits_indptr": torch.arange(bs + 1, dtype=torch.int32, device="cuda"),
"split_logits": torch.empty((bs, 1, N_HEADS, V_DIM), dtype=torch.float32, device="cuda"),
"split_lse": torch.empty((bs, 1, N_HEADS, 1), dtype=torch.float32, device="cuda"),
"out_bf16": torch.empty((bs, N_HEADS, V_DIM), dtype=BF16_T, device="cuda"),
"q_fp8": torch.empty((bs, N_HEADS, QK_DIM), dtype=FP8_T, device="cuda"),
"q_scale": STATIC_Q_SCALE.clone(),
}
# Eager state init for all leaderboard shapes.
SHAPE_STATE = {
shape: _build_shape_state(shape[0], shape[1], cfg[0], cfg[1], cfg[2], cfg[3])
for shape, cfg in SHAPE_POLICY.items()
}
BF16_STATE = {
shape: _build_fwd_a8w8_1k_state(shape[0], shape[1], 2)
for shape in FWD_A8W8_1K_SHAPES
}
NP8K_STATE = {
(bs, 8192): _build_np8k_state(bs, 8192, ps)
for bs, ps in NP_8K_PAGE.items()
}
def _quantize_query_to_fp8(q_bf16: torch.Tensor, q_fp8: torch.Tensor, q_scale: torch.Tensor) -> None:
if HAS_STATIC_QUANT:
static_per_tensor_quant(q_fp8, q_bf16, STATIC_Q_SCALE)
return
if HAS_DYNAMIC_QUANT:
dynamic_per_tensor_quant(q_fp8, q_bf16.view_as(q_fp8), q_scale)
return
finfo = torch.finfo(FP8_T)
amax = q_bf16.abs().amax().clamp(min=1e-12)
scale = amax / finfo.max
q_fp8.copy_((q_bf16 / scale).clamp(finfo.min, finfo.max).to(FP8_T))
q_scale.fill_(scale.item())
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
kv_len = config["kv_seq_len"]
total_q = q.shape[0]
shape = (bs, kv_len)
if shape in FWD_A8W8_1K_SHAPES:
fwd_state = BF16_STATE[shape]
kv_fp8, kv_scale = kv_data["fp8"]
page_size = fwd_state["page_size"]
kv_4d = kv_fp8.view(kv_fp8.shape[0] // page_size, page_size, N_KV_HEADS, QK_DIM)
q_fp8 = fwd_state["q_fp8"]
q_scale = fwd_state["q_scale"]
_quantize_query_to_fp8(q, q_fp8, q_scale)
out = fwd_state["out_bf16"]
mla_decode_fwd(
q=q_fp8,
kv_buffer=kv_4d,
o=out,
qo_indptr=qo_indptr,
kv_indptr=fwd_state["kv_pages_prefix"],
kv_indices=fwd_state["kv_indices"],
kv_last_page_lens=fwd_state["kv_last_page_lens"],
max_seqlen_q=1,
page_size=page_size,
nhead_kv=1,
sm_scale=ATTN_SCALE,
q_scale=q_scale,
kv_scale=kv_scale,
intra_batch_mode=True,
)
return out[:total_q]
if shape in NP8K_STATE:
np_state = NP8K_STATE[shape]
kv_fp8, kv_scale = kv_data["fp8"]
page_size = np_state["page_size"]
kv_4d = kv_fp8.view(kv_fp8.shape[0] // page_size, page_size, N_KV_HEADS, QK_DIM)
q_fp8 = np_state["q_fp8"]
q_scale = np_state["q_scale"]
_quantize_query_to_fp8(q, q_fp8, q_scale)
out = np_state["out_bf16"]
split_logits = np_state["split_logits"]
split_lse = np_state["split_lse"]
out.zero_()
split_lse.fill_(float("-inf"))
mla_decode_stage1_asm_fwd(
q_fp8,
kv_4d,
qo_indptr,
np_state["kv_pages_prefix"],
np_state["kv_indices"],
np_state["kv_last_page_lens"],
np_state["num_kv_splits_indptr"],
None,
None,
None,
1,
page_size,
N_KV_HEADS,
ATTN_SCALE,
split_logits,
split_lse,
out,
q_scale,
kv_scale,
)
return out[:total_q]
state = SHAPE_STATE[shape]
kv_fp8, kv_scale = kv_data["fp8"]
page_size = state["page_size"]
kv_4d = kv_fp8.view(kv_fp8.shape[0] // page_size, page_size, N_KV_HEADS, QK_DIM)
q_fp8 = state["q_fp8"]
q_scale = state["q_scale"]
_quantize_query_to_fp8(q, q_fp8, q_scale)
mla_decode_stage1_asm_fwd(
q_fp8,
kv_4d,
state["qo_prefix"],
state["kv_pages_prefix"],
state["kv_indices"],
state["kv_last_page_len"],
None,
state["wm"],
state["wi"],
state["wis"],
1,
page_size,
N_KV_HEADS,
ATTN_SCALE,
state["partial_logits"],
state["partial_lse"],
state["out_bf16"],
q_scale,
kv_scale,
)
if state["splits"] > 1:
mla_reduce_v1(
state["partial_logits"],
state["partial_lse"],
state["ri"],
state["rfm"],
state["rpm"],
1,
state["out_bf16"],
None,
)
return state["out_bf16"][:total_q]
scrolls · 321 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 687679.
#!POPCORN leaderboard amd-mixed-mla#!POPCORN gpu MI355X- """- Refactored equivalent of my_submission_demo.py.-- Behavior/perf intent:- - identical shape policy and kernel path (a8w8 via stage1_asm + reduce_v1)- - identical eager buffer/materialization strategy- - identical quantization fallback chain- """-+ import osimport torchfrom task import input_t, output_tfrom aiter import dtypes as aiter_dtypesfrom aiter import get_mla_metadata_info_v1, get_mla_metadata_v1from aiter import mla_decode_stage1_asm_fwd, mla_reduce_v1+ from aiter.mla import mla_decode_fwdHAS_STATIC_QUANT = Falsetry:- from aiter.ops.quant import static_per_tensor_quant+ # Prefer JIT quant op when available.+ from aiter.jit.module_quant import static_per_tensor_quantHAS_STATIC_QUANT = True- except ImportError:- pass+ except Exception:+ try:+ from aiter.ops.quant import static_per_tensor_quant+ HAS_STATIC_QUANT = True+ except Exception:+ passtry:from aiter.ops.quant import dynamic_per_tensor_quant⋯ 12 unchanged linesSTATIC_Q_SCALE = torch.tensor([0.1], dtype=torch.float32, device="cuda")- # (batch_size, kv_seq_len) -> (page_size, num_kv_splits, is_sparse)+ # (batch_size, kv_seq_len) -> (page_size, num_kv_splits, is_sparse, kv_granularity_override)SHAPE_POLICY = {- (4, 1024): (1, 8, False),- (4, 8192): (8, 16, False),- (32, 1024): (1, 8, False),- (32, 8192): (8, 12, False),- (64, 1024): (2, 4, False),- (64, 8192): (8, 12, False),- (256, 1024): (2, 1, True),- (256, 8192): (8, 8, False),+ (4, 1024): (2, 8, False, None),+ (4, 8192): (8, 16, False, None),+ (32, 1024): (2, 8, False, None),+ (32, 8192): (8, 12, False, None),+ (64, 1024): (2, 1, False, 4),+ (64, 8192): (8, 12, False, None),+ (256, 1024): (2, 1, True, 2),+ (256, 8192): (8, 8, False, None),}+ NP_8K_PAGE = {+ 4: 2048,+ 32: 1024,+ 64: 512,+ 256: 2048,+ }- def _build_shape_state(bs: int, kv_len: int, page_size: int, splits: int, sparse: bool) -> dict:+ FWD_A8W8_1K_SHAPES = {+ (4, 1024),+ (32, 1024),+ }+++ def _build_shape_state(+ bs: int,+ kv_len: int,+ page_size: int,+ splits: int,+ sparse: bool,+ kv_granularity_override: int | None,+ ) -> dict:pages_per_req = kv_len // page_sizetotal_pages = bs * pages_per_req⋯ 16 unchanged lineswork_meta = [torch.empty(sz, dtype=dt, device="cuda") for sz, dt in meta_info]wm, wi, wis, ri, rfm, rpm = work_meta+ kv_granularity = kv_granularity_override if kv_granularity_override is not None else max(page_size, 16)get_mla_metadata_v1(qo_prefix,kv_pages_prefix,⋯ 8 unchanged linesrfm,rpm,page_size=page_size,- kv_granularity=max(page_size, 16),+ kv_granularity=kv_granularity,max_seqlen_qo=1,uni_seqlen_qo=1,fast_mode=True,⋯ 25 unchanged lines}+ def _build_fwd_a8w8_1k_state(bs: int, kv_len: int, page_size: int = 2) -> dict:+ pages_per_req = kv_len // page_size+ total_pages = bs * pages_per_req+ return {+ "page_size": page_size,+ "kv_indices": torch.arange(total_pages, dtype=torch.int32, device="cuda"),+ "kv_pages_prefix": torch.arange(bs + 1, dtype=torch.int32, device="cuda") * pages_per_req,+ "kv_last_page_lens": torch.full((bs,), page_size, dtype=torch.int32, device="cuda"),+ "out_bf16": torch.empty((bs, N_HEADS, V_DIM), dtype=BF16_T, device="cuda"),+ "q_fp8": torch.empty((bs, N_HEADS, QK_DIM), dtype=FP8_T, device="cuda"),+ "q_scale": STATIC_Q_SCALE.clone(),+ }+++ def _build_np8k_state(bs: int, kv_len: int, page_size: int) -> dict:+ pages_per_req = kv_len // page_size+ total_pages = bs * pages_per_req+ return {+ "page_size": page_size,+ "kv_indices": torch.arange(total_pages, dtype=torch.int32, device="cuda"),+ "kv_pages_prefix": torch.arange(bs + 1, dtype=torch.int32, device="cuda") * pages_per_req,+ "kv_last_page_lens": torch.full((bs,), page_size, dtype=torch.int32, device="cuda"),+ "num_kv_splits_indptr": torch.arange(bs + 1, dtype=torch.int32, device="cuda"),+ "split_logits": torch.empty((bs, 1, N_HEADS, V_DIM), dtype=torch.float32, device="cuda"),+ "split_lse": torch.empty((bs, 1, N_HEADS, 1), dtype=torch.float32, device="cuda"),+ "out_bf16": torch.empty((bs, N_HEADS, V_DIM), dtype=BF16_T, device="cuda"),+ "q_fp8": torch.empty((bs, N_HEADS, QK_DIM), dtype=FP8_T, device="cuda"),+ "q_scale": STATIC_Q_SCALE.clone(),+ }++# Eager state init for all leaderboard shapes.SHAPE_STATE = {- shape: _build_shape_state(shape[0], shape[1], cfg[0], cfg[1], cfg[2])+ shape: _build_shape_state(shape[0], shape[1], cfg[0], cfg[1], cfg[2], cfg[3])for shape, cfg in SHAPE_POLICY.items()}+ BF16_STATE = {+ shape: _build_fwd_a8w8_1k_state(shape[0], shape[1], 2)+ for shape in FWD_A8W8_1K_SHAPES+ }+ NP8K_STATE = {+ (bs, 8192): _build_np8k_state(bs, 8192, ps)+ for bs, ps in NP_8K_PAGE.items()+ }def _quantize_query_to_fp8(q_bf16: torch.Tensor, q_fp8: torch.Tensor, q_scale: torch.Tensor) -> None:⋯ 13 unchanged linesdef custom_kernel(data: input_t) -> output_t:- q, kv_data, _, _, config = data+ q, kv_data, qo_indptr, kv_indptr, config = databs = config["batch_size"]kv_len = config["kv_seq_len"]total_q = q.shape[0]+ shape = (bs, kv_len)- state = SHAPE_STATE[(bs, kv_len)]+ if shape in FWD_A8W8_1K_SHAPES:+ fwd_state = BF16_STATE[shape]+ kv_fp8, kv_scale = kv_data["fp8"]+ page_size = fwd_state["page_size"]+ kv_4d = kv_fp8.view(kv_fp8.shape[0] // page_size, page_size, N_KV_HEADS, QK_DIM)+ q_fp8 = fwd_state["q_fp8"]+ q_scale = fwd_state["q_scale"]+ _quantize_query_to_fp8(q, q_fp8, q_scale)+ out = fwd_state["out_bf16"]+ mla_decode_fwd(+ q=q_fp8,+ kv_buffer=kv_4d,+ o=out,+ qo_indptr=qo_indptr,+ kv_indptr=fwd_state["kv_pages_prefix"],+ kv_indices=fwd_state["kv_indices"],+ kv_last_page_lens=fwd_state["kv_last_page_lens"],+ max_seqlen_q=1,+ page_size=page_size,+ nhead_kv=1,+ sm_scale=ATTN_SCALE,+ q_scale=q_scale,+ kv_scale=kv_scale,+ intra_batch_mode=True,+ )+ return out[:total_q]++ if shape in NP8K_STATE:+ np_state = NP8K_STATE[shape]+ kv_fp8, kv_scale = kv_data["fp8"]+ page_size = np_state["page_size"]+ kv_4d = kv_fp8.view(kv_fp8.shape[0] // page_size, page_size, N_KV_HEADS, QK_DIM)+ q_fp8 = np_state["q_fp8"]+ q_scale = np_state["q_scale"]+ _quantize_query_to_fp8(q, q_fp8, q_scale)++ out = np_state["out_bf16"]+ split_logits = np_state["split_logits"]+ split_lse = np_state["split_lse"]+ out.zero_()+ split_lse.fill_(float("-inf"))+ mla_decode_stage1_asm_fwd(+ q_fp8,+ kv_4d,+ qo_indptr,+ np_state["kv_pages_prefix"],+ np_state["kv_indices"],+ np_state["kv_last_page_lens"],+ np_state["num_kv_splits_indptr"],+ None,+ None,+ None,+ 1,+ page_size,+ N_KV_HEADS,+ ATTN_SCALE,+ split_logits,+ split_lse,+ out,+ q_scale,+ kv_scale,+ )+ return out[:total_q]++ state = SHAPE_STATE[shape]kv_fp8, kv_scale = kv_data["fp8"]page_size = state["page_size"]kv_4d = kv_fp8.view(kv_fp8.shape[0] // page_size, page_size, N_KV_HEADS, QK_DIM)
scrolls · 237 diff lines total
Best evidence level for this revision: reported
JSON