submission 687679
NinoHeather · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 199 lines, June 9 Researcher Reciprocity License v1.0.
my_submission_refact.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-687679?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:76502f69b76754b245799313638f0cce4868a469e7f360067410d02f420c47ae
license declaredunknown
license concludedunknown
authorsNinoHeather
imported2026-08-15
Kernel source
my_submission_refact.py199 lines
#!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 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
HAS_STATIC_QUANT = False
try:
from aiter.ops.quant import static_per_tensor_quant
HAS_STATIC_QUANT = True
except ImportError:
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)
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),
}
def _build_shape_state(bs: int, kv_len: int, page_size: int, splits: int, sparse: bool) -> 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
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=max(page_size, 16),
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(),
}
# Eager state init for all leaderboard shapes.
SHAPE_STATE = {
shape: _build_shape_state(shape[0], shape[1], cfg[0], cfg[1], cfg[2])
for shape, cfg in SHAPE_POLICY.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, _, _, config = data
bs = config["batch_size"]
kv_len = config["kv_seq_len"]
total_q = q.shape[0]
state = SHAPE_STATE[(bs, kv_len)]
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 · 199 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 683428.
⋯ 1 unchanged lines#!POPCORN gpu MI355X"""- Two-stage MLA decode (stage1_asm + reduce_v1) with fp8 Q static quantization.- Per-shape tuned page_size / num_kv_splits. Eager buffer initialization.- All shapes use a8w8 path (fp8 Q + fp8 KV) for maximum throughput.+ 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 torch⋯ 2 unchanged linesfrom aiter import get_mla_metadata_info_v1, get_mla_metadata_v1from aiter import mla_decode_stage1_asm_fwd, mla_reduce_v1- _USE_STATIC = False++ HAS_STATIC_QUANT = Falsetry:from aiter.ops.quant import static_per_tensor_quant- _USE_STATIC = True+ HAS_STATIC_QUANT = Trueexcept ImportError:pass+try:from aiter.ops.quant import dynamic_per_tensor_quant- _HAS_DYN = True+ HAS_DYNAMIC_QUANT = Trueexcept ImportError:- _HAS_DYN = False+ HAS_DYNAMIC_QUANT = False- FP8 = aiter_dtypes.fp8- BF16 = torch.bfloat16- NH = 16- NKV = 1- DQ = 576- DV = 512- SM = 1.0 / (DQ ** 0.5)- _STATIC_SCALE = torch.tensor([0.1], dtype=torch.float32, device='cuda')- # (page_size, num_splits, is_sparse)- # ps: larger pages amortize page-table overhead for long kv- # ns: fewer splits = less reduce overhead; more splits = better kv parallelism- # sparse: helps with large batch + few splits metadata layout- _CFG = {- (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),+ 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)+ 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),}- _meta = {}- for (_bs, _kv), (_ps, _ns, _sp) in _CFG.items():- _npb = _kv // _ps- _tp = _bs * _npb+ def _build_shape_state(bs: int, kv_len: int, page_size: int, splits: int, sparse: bool) -> dict:+ pages_per_req = kv_len // page_size+ total_pages = bs * pages_per_req- _qo = torch.arange(_bs + 1, dtype=torch.int32, device='cuda')- _kip = torch.arange(_bs + 1, dtype=torch.int32, device='cuda') * _npb- _klp = torch.full((_bs,), _ps, dtype=torch.int32, device='cuda')- _ki = torch.arange(_tp, dtype=torch.int32, device='cuda')+ 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")- _info = get_mla_metadata_info_v1(- _bs, 1, NH, FP8, FP8,- is_sparse=_sp, fast_mode=True,- num_kv_splits=_ns, intra_batch_mode=True)- _w = [torch.empty(s, dtype=t, device='cuda') for s, t in _info]- _wm, _wi, _wis, _ri, _rfm, _rpm = _w+ 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_metaget_mla_metadata_v1(- _qo, _kip, _klp,- NH // NKV, NKV, False,- _wm, _wis, _wi, _ri, _rfm, _rpm,- page_size=_ps, kv_granularity=max(_ps, 16),- max_seqlen_qo=1, uni_seqlen_qo=1,- fast_mode=True, max_split_per_batch=_ns,- intra_batch_mode=True, dtype_q=FP8, dtype_kv=FP8)+ 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=max(page_size, 16),+ 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,+ )- _np = _rpm.size(0)- _meta[(_bs, _kv)] = {- 'qo': _qo, 'kip': _kip, 'klp': _klp, 'ki': _ki,- 'wm': _wm, 'wi': _wi, 'wis': _wis,- 'ri': _ri, 'rfm': _rfm, 'rpm': _rpm,- 'logits': torch.empty((_np, 1, NH, DV), dtype=torch.float32, device='cuda'),- 'lse': torch.empty((_np, 1, NH, 1), dtype=torch.float32, device='cuda'),- 'out': torch.empty((_bs, NH, DV), dtype=BF16, device='cuda'),- 'qbuf': torch.empty((_bs, NH, DQ), dtype=FP8, device='cuda'),- 'qscale': _STATIC_SCALE.clone(),- 'ps': _ps, 'ns': _ns,+ 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(),}+ # Eager state init for all leaderboard shapes.+ SHAPE_STATE = {+ shape: _build_shape_state(shape[0], shape[1], cfg[0], cfg[1], cfg[2])+ for shape, cfg in SHAPE_POLICY.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, _, _, config = data- bs = config['batch_size']- kv = config['kv_seq_len']- tq = q.shape[0]+ bs = config["batch_size"]+ kv_len = config["kv_seq_len"]+ total_q = q.shape[0]- m = _meta[(bs, kv)]- kvf, kvs = kv_data['fp8']- ps = m['ps']- kv4 = kvf.view(kvf.shape[0] // ps, ps, NKV, DQ)+ state = SHAPE_STATE[(bs, kv_len)]+ 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)- qbuf = m['qbuf']- qsc = m['qscale']- if _USE_STATIC:- static_per_tensor_quant(qbuf, q, _STATIC_SCALE)- elif _HAS_DYN:- dynamic_per_tensor_quant(qbuf, q.view_as(qbuf), qsc)- else:- finfo = torch.finfo(FP8)- amax = q.abs().amax().clamp(min=1e-12)- scale = amax / finfo.max- qbuf.copy_((q / scale).clamp(finfo.min, finfo.max).to(FP8))- qsc.fill_(scale.item())+ q_fp8 = state["q_fp8"]+ q_scale = state["q_scale"]+ _quantize_query_to_fp8(q, q_fp8, q_scale)mla_decode_stage1_asm_fwd(- qbuf, kv4,- m['qo'], m['kip'], m['ki'], m['klp'],+ q_fp8,+ kv_4d,+ state["qo_prefix"],+ state["kv_pages_prefix"],+ state["kv_indices"],+ state["kv_last_page_len"],None,- m['wm'], m['wi'], m['wis'],- 1, ps, NKV, SM,- m['logits'], m['lse'], m['out'],- qsc, kvs)+ 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 m['ns'] > 1:+ if state["splits"] > 1:mla_reduce_v1(- m['logits'], m['lse'],- m['ri'], m['rfm'], m['rpm'],- 1, m['out'], None)+ state["partial_logits"],+ state["partial_lse"],+ state["ri"],+ state["rfm"],+ state["rpm"],+ 1,+ state["out_bf16"],+ None,+ )- return m['out'][:tq]+ return state["out_bf16"][:total_q]
scrolls · 287 diff lines total
Best evidence level for this revision: reported
JSON