submission 697219
parcadei · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 327 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-697219?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:52d89d34c1e6a72bea2e93c2d3df1b6d9775d9b10a0515155fa309b3f0535a95
license declaredunknown
license concludedunknown
authorsparcadei
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
- S2: Persistent FP8 (gran=128, ns=32) + C++ reduceKernel source
submission.py327 lines
"""
Mixed-MLA submission: aiter ASM with hybrid dispatch.
Architecture:
- S1: BF16 16-split + C++ reduce (no Q copy, no FP8 overhead)
- S2: Persistent FP8 (gran=128, ns=32) + C++ reduce
- S3: BF16 4-split + C++ reduce
- S4: FP8 8-split + C++ reduce
- S5: Persistent FP8 (gran=64, ns=32) + C++ reduce
- S6: FP8 4-split + C++ reduce
- S7/S8: NP1 FP8 (1-split, no reduce)
- Closure-based dispatch: pre-built callables per shape capture all tensor refs
"""
from __future__ import annotations
import math
import os
from typing import Any
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
import torch
import torch.nn.functional as F
from task import input_t, output_t
PAGE_SIZE = 1
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / math.sqrt(QK_HEAD_DIM)
# ---- Aiter lazy-load --------------------------------------------------------
_AITER_API: dict[str, Any] | None = None
_AITER_IMPORT_ERROR: Exception | None = None
_FN_STAGE1: Any = None
_FN_METADATA: Any = None
_FN_REDUCE: Any = None
_FP8_DTYPE: Any = None
_META_INFO_FN: Any = None
def _load_aiter() -> dict[str, Any]:
global _AITER_API, _AITER_IMPORT_ERROR, _FN_STAGE1, _FN_METADATA, _FN_REDUCE, _FP8_DTYPE, _META_INFO_FN
if _AITER_API is not None:
return _AITER_API
if _AITER_IMPORT_ERROR is not None:
raise _AITER_IMPORT_ERROR
try:
import aiter
from aiter import dtypes as aiter_dtypes
from aiter.ops.attention import get_mla_metadata_info_v1
_AITER_API = {"ok": True}
_FP8_DTYPE = aiter_dtypes.fp8
_FN_STAGE1 = aiter.mla_decode_stage1_asm_fwd
_FN_METADATA = aiter.get_mla_metadata_v1
_FN_REDUCE = aiter.mla_reduce_v1
_META_INFO_FN = get_mla_metadata_info_v1
return _AITER_API
except Exception as exc:
_AITER_IMPORT_ERROR = exc
raise
def _torch_fallback(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
kv_bf16 = kv_data["bf16"]
sm_scale = float(config.get("sm_scale", SM_SCALE))
kv_lora_rank = int(config["kv_lora_rank"])
out_chunks = []
batch_size = qo_indptr.shape[0] - 1
for batch_idx in range(batch_size):
q_start = int(qo_indptr[batch_idx].item())
q_end = int(qo_indptr[batch_idx + 1].item())
kv_start = int(kv_indptr[batch_idx].item())
kv_end = int(kv_indptr[batch_idx + 1].item())
qi = q[q_start:q_end].float().permute(1, 0, 2)
kv_rows = kv_bf16[kv_start:kv_end, 0].float()
scores = torch.matmul(qi * sm_scale, kv_rows.T)
probs = F.softmax(scores, dim=-1)
values = kv_rows[:, :kv_lora_rank]
out = torch.matmul(probs, values).permute(1, 0, 2)
out_chunks.append(out.to(torch.bfloat16))
return torch.cat(out_chunks, dim=0)
# ---- Entry point ------------------------------------------------------------
_SHAPE_DISPATCH: dict[tuple, Any] = {}
# Shape -> num_splits for FP8 multi-split + C++ reduce
_FP8_SPLITS: dict[tuple[int, int], int] = {
(32, 8192): 8, # S4: FP8 8-split + C++ reduce
(64, 8192): 4, # S6: FP8 4-split + C++ reduce
}
# Shapes using NP1 (1-split, no reduce, stage1 writes directly to output)
_NP1_SHAPES: set[tuple[int, int]] = {
(256, 1024), # S7
(256, 8192), # S8
}
# Shape -> num_splits for BF16 multi-split + C++ reduce
_BF16_SPLITS: dict[tuple[int, int], int] = {
(4, 1024): 16, # S1: BF16 16-split + C++ reduce
(32, 1024): 4, # S3: BF16 4-split — FP8 non-persistent fails ranked
}
# Shape -> (kv_granularity, num_splits) for persistent mode + C++ reduce
_PERSISTENT_CONFIG: dict[tuple[int, int], tuple[int, int]] = {
(4, 8192): (128, 32), # S2: persistent, gran=128
(64, 1024): (64, 32), # S5: persistent, gran=64
}
def _build_dispatch(key, q, data):
"""Cold path: pre-allocate tensors and build a closure for the hot path."""
if _FN_STAGE1 is None:
_load_aiter()
config = data[4]
bs, kv_len = key
nh = int(config["num_heads"])
qs = int(config["q_seq_len"])
total_kv = bs * kv_len
device = q.device
kv_fp8 = data[1]["fp8"][0]
kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)
output = torch.empty((q.shape[0], nh, V_HEAD_DIM), dtype=torch.bfloat16, device=device)
rows = q.numel() // q.shape[-1]
q_flat_shape = (rows, q.shape[-1])
q_fp8_buf = torch.empty(q_flat_shape, dtype=_FP8_DTYPE, device=device)
q_shaped = q_fp8_buf.view(-1, nh, QK_HEAD_DIM)
# Cached unit scale for cast-only FP8 quantization.
# Persistent ASM kernel does NOT modify q_scale in-place (verified: reference
# caches a single unit_scale tensor across calls without resetting).
# Non-persistent ASM kernel DOES modify q_scale — use fill_(1.0) there.
q_scale = torch.ones(1, dtype=torch.float32, device=device)
qo_indptr = data[2].clone()
kv_indptr = data[3].clone()
kv_view_shape = (kv_fp8.shape[0], PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)
bf16_splits = _BF16_SPLITS.get(key)
persistent_config = _PERSISTENT_CONFIG.get(key)
fp8_splits = _FP8_SPLITS.get(key)
is_np1 = key in _NP1_SHAPES
if bf16_splits is not None:
# BF16 multi-split + C++ reduce (S1, S3)
total_q = q.shape[0]
num_splits = bf16_splits
splits_indptr = torch.arange(0, (bs + 1) * num_splits, num_splits, dtype=torch.int, device=device)
kv_last = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
np_logits = torch.empty((total_q, num_splits, nh, V_HEAD_DIM), dtype=torch.float32, device=device)
np_lse = torch.empty((total_q, num_splits, nh, 1), dtype=torch.float32, device=device)
kv_bf16 = data[1]["bf16"]
kv_bf16_view = (kv_bf16.shape[0], PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)
np_total = total_q * num_splits
reduce_indptr = torch.arange(0, np_total + num_splits, num_splits, dtype=torch.int32, device=device)[:total_q + 1]
reduce_partial_map = torch.arange(np_total, dtype=torch.int32, device=device)
rl_view = np_logits.reshape(np_total, 1, nh, V_HEAD_DIM)
rls_view = np_lse.reshape(np_total, 1, nh, 1)
def _dispatch_bf16(data):
q_bf16 = data[0].view(-1, nh, QK_HEAD_DIM)
kv_v = data[1]["bf16"].view(kv_bf16_view)
_FN_STAGE1(
q_bf16, kv_v,
qo_indptr, kv_indptr, kv_indices,
kv_last, splits_indptr, None, None, None,
qs, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
np_logits, np_lse, output, None, None,
)
_FN_REDUCE(rl_view, rls_view, reduce_indptr, None, reduce_partial_map, qs, output, None)
return output
_SHAPE_DISPATCH[key] = _dispatch_bf16
return _dispatch_bf16
elif persistent_config is not None:
# PERSISTENT mode: metadata + persistent ASM kernel + C++ reduce (S2, S5)
kv_gran, num_splits = persistent_config
total_q = q.shape[0]
kv_last = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
meta_info = _META_INFO_FN(
batch_size=bs,
max_seqlen_qo=qs,
num_head_qo=nh,
q_dtype=_FP8_DTYPE,
kv_dtype=_FP8_DTYPE,
is_sparse=False,
fast_mode=True,
num_kv_splits=num_splits,
intra_batch_mode=True,
)
work_metadata_ptrs = torch.empty(meta_info[0][0], dtype=meta_info[0][1], device=device)
work_indptr = torch.empty(meta_info[1][0], dtype=meta_info[1][1], device=device)
work_info_set = torch.empty(meta_info[2][0], dtype=meta_info[2][1], device=device)
reduce_indptr = torch.empty(meta_info[3][0], dtype=meta_info[3][1], device=device)
reduce_final_map = torch.empty(meta_info[4][0], dtype=meta_info[4][1], device=device)
reduce_partial_map = torch.empty(meta_info[5][0], dtype=meta_info[5][1], device=device)
_FN_METADATA(
qo_indptr, kv_indptr, kv_last,
nh, NUM_KV_HEADS, True,
work_metadata_ptrs, work_info_set, work_indptr,
reduce_indptr, reduce_final_map, reduce_partial_map,
PAGE_SIZE, kv_gran,
qs, qs, True, -1, num_splits, True,
_FP8_DTYPE, _FP8_DTYPE,
)
n_partials = reduce_partial_map.shape[0]
logits = torch.empty(
(n_partials * qs, 1, nh, V_HEAD_DIM),
dtype=torch.float32, device=device,
)
attn_lse = torch.empty(
(n_partials * qs, 1, nh, 1),
dtype=torch.float32, device=device,
)
def _dispatch_persistent(data):
q_fp8_buf.copy_(data[0].view(q_flat_shape))
# No fill_(1.0) needed: persistent ASM kernel does not modify q_scale
fp8 = data[1]["fp8"]
kv_v = fp8[0].view(kv_view_shape)
_FN_STAGE1(
q_shaped, kv_v,
qo_indptr, kv_indptr, kv_indices,
kv_last, None,
work_metadata_ptrs, work_indptr, work_info_set,
qs, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
logits, attn_lse, output, q_scale, fp8[1],
)
_FN_REDUCE(
logits, attn_lse,
reduce_indptr, reduce_final_map, reduce_partial_map,
qs, output, None,
)
return output
_SHAPE_DISPATCH[key] = _dispatch_persistent
return _dispatch_persistent
elif fp8_splits is not None:
# FP8 multi-split + C++ reduce (S4, S6)
total_q = q.shape[0]
num_splits = fp8_splits
splits_indptr = torch.arange(0, (bs + 1) * num_splits, num_splits, dtype=torch.int, device=device)
kv_last = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
np_logits = torch.empty((total_q, num_splits, nh, V_HEAD_DIM), dtype=torch.float32, device=device)
np_lse = torch.empty((total_q, num_splits, nh, 1), dtype=torch.float32, device=device)
np_total = total_q * num_splits
reduce_indptr_fp8 = torch.arange(0, np_total + num_splits, num_splits, dtype=torch.int32, device=device)[:total_q + 1]
reduce_partial_map_fp8 = torch.arange(np_total, dtype=torch.int32, device=device)
rl_view = np_logits.reshape(np_total, 1, nh, V_HEAD_DIM)
rls_view = np_lse.reshape(np_total, 1, nh, 1)
def _dispatch_fp8(data):
q_fp8_buf.copy_(data[0].view(q_flat_shape))
fp8 = data[1]["fp8"]
kv_v = fp8[0].view(kv_view_shape)
_FN_STAGE1(
q_shaped, kv_v,
qo_indptr, kv_indptr, kv_indices,
kv_last, splits_indptr, None, None, None,
qs, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
np_logits, np_lse, output, q_scale, fp8[1],
)
_FN_REDUCE(rl_view, rls_view, reduce_indptr_fp8, None, reduce_partial_map_fp8, qs, output, None)
return output
_SHAPE_DISPATCH[key] = _dispatch_fp8
return _dispatch_fp8
elif is_np1:
# NP1: 1-split, no reduce needed, stage1 writes bf16 directly to output
total_q = q.shape[0]
indptr = torch.arange(0, bs + 1, dtype=torch.int, device=device)
kv_last = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
logits = output.view(total_q, 1, nh, V_HEAD_DIM)
attn_lse = torch.empty((total_q, 1, nh, 1), dtype=torch.float32, device=device)
def _dispatch_np1(data):
q_fp8_buf.copy_(data[0].view(q_flat_shape))
fp8 = data[1]["fp8"]
kv_v = fp8[0].view(kv_view_shape)
_FN_STAGE1(
q_shaped, kv_v,
qo_indptr, kv_indptr, kv_indices,
kv_last, indptr, None, None, None,
qs, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
logits, attn_lse, output, q_scale, fp8[1],
)
return output
_SHAPE_DISPATCH[key] = _dispatch_np1
return _dispatch_np1
else:
raise RuntimeError(f"Shape {key} not in dispatch table")
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
q = data[0]
if not q.is_cuda:
return _torch_fallback(data)
config = data[4]
key = (int(config["batch_size"]), int(config["kv_seq_len"]))
dispatch = _SHAPE_DISPATCH.get(key)
if dispatch is not None:
return dispatch(data)
dispatch = _build_dispatch(key, q, data)
return dispatch(data)
scrolls · 327 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