submission 670551
南北绿豆糕 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 687 lines, June 9 Researcher Reciprocity License v1.0.
submission_0330_2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-670551?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:dba16a52fb65b85f75da7e1ccbb08a16f13a34d6bb7221d4b744f74f4761c46e
license declaredunknown
license concludedunknown
authors南北绿豆糕
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
mla_decode_fwd(*args, persistent=True, **kwargs)Kernel source
submission_0330_2.py687 lines
import os
import sys
import torch
from task import input_t, output_t
NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
PAGE_SIZE = 1
PG2_PAGE_SIZE = 2
NUM_KV_SPLITS = 64
KV_GRANULARITY = max(PAGE_SIZE, 16)
PG2_KV_GRANULARITY = 16 # Toggle to 64 to test the alternate leaderboard hint.
SM_SCALE = QK_HEAD_DIM ** -0.5
_DEVNULL_HANDLE = None
def _mute_process_stderr_forever() -> None:
global _DEVNULL_HANDLE
try:
sys.stderr.flush()
except Exception:
pass
try:
devnull_fd = os.open(os.devnull, os.O_WRONLY)
os.dup2(devnull_fd, 2)
_DEVNULL_HANDLE = os.fdopen(devnull_fd, "w")
except Exception:
return
try:
sys.stderr = _DEVNULL_HANDLE
except Exception:
pass
_mute_process_stderr_forever()
try:
from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
FP8_DTYPE = aiter_dtypes.fp8
_AITER_AVAILABLE = True
except Exception: # pragma: no cover - runtime availability depends on remote env
mla_decode_fwd = None
aiter_dtypes = None
get_mla_metadata_info_v1 = None
get_mla_metadata_v1 = None
FP8_DTYPE = getattr(torch, "float8_e4m3fnuz", torch.uint8)
_AITER_AVAILABLE = False
class _Meta:
__slots__ = (
"work_meta_data",
"work_indptr",
"work_info_set",
"reduce_indptr",
"reduce_final_map",
"reduce_partial_map",
)
class _PagedCache:
__slots__ = (
"page_size",
"kv_granularity",
"kv_indptr_i32",
"kv_last_page_len",
"kv_indices",
"num_pages",
"meta_bf16_fp8",
"meta_fp8_fp8",
)
class _CaseState:
__slots__ = (
"batch_size",
"q_seq_len",
"kv_seq_len",
"device",
"device_key",
"total_q",
"total_kv",
"prefer_a8w8",
"dispatch_mode",
"dispatch_page_size",
"a16w8_use_unit_scale",
"qo_indptr_i32",
"page1",
"page2",
"active_page",
"out",
"unit_scale",
"q_raw_ptr",
"q_bf16",
"q_fp8_src_ptr",
"q_fp8",
"q_scale",
"kv_raw_ptr",
"kv_fp8",
"kv_buffer_4d",
"kv_scale_raw_ptr",
"kv_scale",
)
_KV_INDICES_CACHE: dict[tuple[int, tuple[str, int | None]], torch.Tensor] = {}
_STATE_CACHE: dict[tuple[int, int, tuple[str, int | None]], _CaseState] = {}
_LAST_CASE_KEY: tuple[int, int, tuple[str, int | None]] | None = None
_LAST_STATE: _CaseState | None = None
def _device_key(device: torch.device) -> tuple[str, int | None]:
return (device.type, device.index)
def _make_uniform_indptr(batch_size: int, stride: int, device: torch.device) -> torch.Tensor:
return torch.arange(batch_size + 1, dtype=torch.int32, device=device).mul_(stride)
def _get_cached_kv_indices(total_kv_len: int, device: torch.device) -> torch.Tensor:
key = (total_kv_len, _device_key(device))
cached = _KV_INDICES_CACHE.get(key)
if cached is not None:
return cached
kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device=device)
_KV_INDICES_CACHE[key] = kv_indices
return kv_indices
def _build_paged_cache(
batch_size: int,
q_seq_len: int,
kv_seq_len: int,
total_kv: int,
device: torch.device,
page_size: int,
kv_granularity: int,
) -> _PagedCache:
page = _PagedCache()
page.page_size = page_size
page.kv_granularity = kv_granularity
if page_size == 1:
page_count_per_batch = kv_seq_len
page.kv_indptr_i32 = _make_uniform_indptr(batch_size, page_count_per_batch, device)
page.kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)
page.kv_indices = _get_cached_kv_indices(total_kv, device)
page.num_pages = total_kv
else:
page_count_per_batch = (kv_seq_len + page_size - 1) // page_size
last_page_len = kv_seq_len % page_size
if last_page_len == 0:
last_page_len = page_size
page.kv_indptr_i32 = _make_uniform_indptr(batch_size, page_count_per_batch, device)
page.kv_last_page_len = torch.full((batch_size,), last_page_len, dtype=torch.int32, device=device)
page.kv_indices = _get_cached_kv_indices(total_kv, device)[::page_size].div(page_size, rounding_mode="floor")
page.num_pages = batch_size * page_count_per_batch
if _AITER_AVAILABLE:
page.meta_bf16_fp8 = _make_mla_decode_metadata(
batch_size=batch_size,
max_q_len=q_seq_len,
q_dtype=torch.bfloat16,
kv_dtype=FP8_DTYPE,
qo_indptr=_make_uniform_indptr(batch_size, q_seq_len, device),
kv_indptr=page.kv_indptr_i32,
kv_last_page_len=page.kv_last_page_len,
page_size=page.page_size,
kv_granularity=page.kv_granularity,
)
page.meta_fp8_fp8 = _make_mla_decode_metadata(
batch_size=batch_size,
max_q_len=q_seq_len,
q_dtype=FP8_DTYPE,
kv_dtype=FP8_DTYPE,
qo_indptr=_make_uniform_indptr(batch_size, q_seq_len, device),
kv_indptr=page.kv_indptr_i32,
kv_last_page_len=page.kv_last_page_len,
page_size=page.page_size,
kv_granularity=page.kv_granularity,
)
else: # pragma: no cover - metadata only matters in remote runtime
page.meta_bf16_fp8 = None
page.meta_fp8_fp8 = None
return page
def _make_mla_decode_metadata(
batch_size: int,
max_q_len: int,
q_dtype: torch.dtype,
kv_dtype: torch.dtype,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
kv_last_page_len: torch.Tensor,
page_size: int,
kv_granularity: int,
) -> _Meta:
info = get_mla_metadata_info_v1(
batch_size,
max_q_len,
NUM_HEADS,
q_dtype,
kv_dtype,
is_sparse=False,
fast_mode=True,
num_kv_splits=NUM_KV_SPLITS,
intra_batch_mode=True,
)
work = [torch.empty(shape, dtype=dtype, device=qo_indptr.device) for shape, dtype in info]
meta = _Meta()
(
meta.work_meta_data,
meta.work_indptr,
meta.work_info_set,
meta.reduce_indptr,
meta.reduce_final_map,
meta.reduce_partial_map,
) = work
get_mla_metadata_v1(
qo_indptr,
kv_indptr,
kv_last_page_len,
NUM_HEADS // NUM_KV_HEADS,
NUM_KV_HEADS,
True,
meta.work_meta_data,
meta.work_info_set,
meta.work_indptr,
meta.reduce_indptr,
meta.reduce_final_map,
meta.reduce_partial_map,
page_size=page_size,
kv_granularity=kv_granularity,
max_seqlen_qo=max_q_len,
uni_seqlen_qo=max_q_len,
fast_mode=True,
max_split_per_batch=NUM_KV_SPLITS,
intra_batch_mode=True,
dtype_q=q_dtype,
dtype_kv=kv_dtype,
)
return meta
def _new_case_state(batch_size: int, q_seq_len: int, kv_seq_len: int, device: torch.device) -> _CaseState:
state = _CaseState()
state.batch_size = batch_size
state.q_seq_len = q_seq_len
state.kv_seq_len = kv_seq_len
state.device = device
state.device_key = _device_key(device)
state.total_q = batch_size * q_seq_len
state.total_kv = batch_size * kv_seq_len
state.prefer_a8w8 = batch_size >= 256 and kv_seq_len >= 8192
state.dispatch_mode = 0
state.dispatch_page_size = 0
state.a16w8_use_unit_scale = False
state.qo_indptr_i32 = _make_uniform_indptr(batch_size, q_seq_len, device)
state.page1 = _build_paged_cache(
batch_size=batch_size,
q_seq_len=q_seq_len,
kv_seq_len=kv_seq_len,
total_kv=state.total_kv,
device=device,
page_size=PAGE_SIZE,
kv_granularity=KV_GRANULARITY,
)
state.page2 = None
if kv_seq_len % PG2_PAGE_SIZE == 0 and state.total_kv % PG2_PAGE_SIZE == 0:
state.page2 = _build_paged_cache(
batch_size=batch_size,
q_seq_len=q_seq_len,
kv_seq_len=kv_seq_len,
total_kv=state.total_kv,
device=device,
page_size=PG2_PAGE_SIZE,
kv_granularity=max(PG2_PAGE_SIZE, PG2_KV_GRANULARITY),
)
state.active_page = state.page1
state.out = torch.empty((state.total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device)
state.unit_scale = torch.ones(1, dtype=torch.float32, device=device)
state.q_raw_ptr = 0
state.q_bf16 = None
state.q_fp8_src_ptr = 0
state.q_fp8 = None
state.q_scale = None
state.kv_raw_ptr = 0
state.kv_fp8 = None
state.kv_buffer_4d = None
state.kv_scale_raw_ptr = 0
state.kv_scale = None
return state
def _get_case_state(batch_size: int, q_seq_len: int, kv_seq_len: int, device: torch.device) -> _CaseState:
global _LAST_CASE_KEY, _LAST_STATE
key = (batch_size, kv_seq_len, _device_key(device))
state = _LAST_STATE
if state is not None and _LAST_CASE_KEY == key:
if state.q_seq_len == q_seq_len:
return state
state = _STATE_CACHE.get(key)
if state is None or state.q_seq_len != q_seq_len:
state = _new_case_state(batch_size, q_seq_len, kv_seq_len, device)
_STATE_CACHE[key] = state
_LAST_CASE_KEY = key
_LAST_STATE = state
return state
def _normalize_bf16_tensor(tensor: torch.Tensor, device: torch.device) -> torch.Tensor:
if tensor.device == device and tensor.dtype == torch.bfloat16 and tensor.is_contiguous():
return tensor
return tensor.to(device=device, dtype=torch.bfloat16).contiguous()
def _normalize_tensor(tensor: torch.Tensor, device: torch.device, dtype: torch.dtype) -> torch.Tensor:
if tensor.device == device and tensor.dtype == dtype and tensor.is_contiguous():
return tensor
return tensor.to(device=device, dtype=dtype).contiguous()
def _normalize_scale(scale: torch.Tensor, device: torch.device) -> torch.Tensor:
if scale.device == device and scale.dtype == torch.float32 and scale.is_contiguous() and scale.numel() == 1:
return scale.reshape(1)
return scale.to(device=device, dtype=torch.float32).reshape(1).contiguous()
def _can_zero_copy_pg2_view(state: _CaseState, kv_fp8: torch.Tensor) -> bool:
if state.page2 is None:
return False
if kv_fp8.dim() != 3 or kv_fp8.shape != (state.total_kv, NUM_KV_HEADS, QK_HEAD_DIM):
return False
if kv_fp8.storage_offset() != 0 or not kv_fp8.is_contiguous():
return False
try:
kv_fp8.view(state.page2.num_pages, PG2_PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)
except RuntimeError:
return False
return True
def _bind_inputs(state: _CaseState, q: torch.Tensor, kv_fp8: torch.Tensor, kv_scale: torch.Tensor) -> None:
q_ptr = q.data_ptr()
if q_ptr != state.q_raw_ptr:
state.q_raw_ptr = q_ptr
state.q_bf16 = _normalize_bf16_tensor(q, state.device)
state.q_fp8_src_ptr = 0
state.q_fp8 = None
state.q_scale = None
kv_ptr = kv_fp8.data_ptr()
if kv_ptr != state.kv_raw_ptr:
state.kv_raw_ptr = kv_ptr
prev_page_size = state.active_page.page_size
use_pg2 = _can_zero_copy_pg2_view(state, kv_fp8)
state.active_page = state.page2 if use_pg2 else state.page1
state.kv_fp8 = _normalize_tensor(kv_fp8, state.device, kv_fp8.dtype)
state.kv_buffer_4d = state.kv_fp8.view(
state.active_page.num_pages,
state.active_page.page_size,
NUM_KV_HEADS,
QK_HEAD_DIM,
)
if state.active_page.page_size != prev_page_size:
state.dispatch_mode = 0
state.dispatch_page_size = 0
kv_scale_ptr = kv_scale.data_ptr()
if kv_scale_ptr != state.kv_scale_raw_ptr:
state.kv_scale_raw_ptr = kv_scale_ptr
state.kv_scale = _normalize_scale(kv_scale, state.device)
def _quantize_fp8_per_tensor(tensor: torch.Tensor, fp8_dtype: torch.dtype) -> tuple[torch.Tensor, torch.Tensor]:
finfo = torch.finfo(fp8_dtype)
amax = tensor.abs().amax().clamp(min=1e-12)
scale = (amax / finfo.max).to(torch.float32).reshape(1)
fp8_tensor = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(fp8_dtype)
if not fp8_tensor.is_contiguous():
fp8_tensor = fp8_tensor.contiguous()
return fp8_tensor, scale
def _ensure_q_fp8(state: _CaseState) -> None:
src_ptr = state.q_bf16.data_ptr()
if src_ptr == state.q_fp8_src_ptr:
return
state.q_fp8_src_ptr = src_ptr
state.q_fp8, state.q_scale = _quantize_fp8_per_tensor(state.q_bf16, FP8_DTYPE)
def _call_mla_decode_fwd(*args, **kwargs) -> None:
try:
mla_decode_fwd(*args, persistent=True, **kwargs)
return
except TypeError as exc:
if "persistent" not in str(exc):
raise
mla_decode_fwd(*args, **kwargs)
def _run_aiter_decode(state: _CaseState, q_input: torch.Tensor, q_scale: torch.Tensor | None, meta: _Meta) -> torch.Tensor:
page = state.active_page
_call_mla_decode_fwd(
q_input,
state.kv_buffer_4d,
state.out,
state.qo_indptr_i32,
page.kv_indptr_i32,
page.kv_indices,
page.kv_last_page_len,
state.q_seq_len,
page_size=page.page_size,
nhead_kv=NUM_KV_HEADS,
sm_scale=SM_SCALE,
logit_cap=0.0,
num_kv_splits=NUM_KV_SPLITS,
q_scale=q_scale,
kv_scale=state.kv_scale,
intra_batch_mode=True,
work_meta_data=meta.work_meta_data,
work_indptr=meta.work_indptr,
work_info_set=meta.work_info_set,
reduce_indptr=meta.reduce_indptr,
reduce_final_map=meta.reduce_final_map,
reduce_partial_map=meta.reduce_partial_map,
)
return state.out
def _run_a16w8(state: _CaseState) -> torch.Tensor:
q_scale = state.unit_scale if state.a16w8_use_unit_scale else None
return _run_aiter_decode(state, state.q_bf16, q_scale, state.active_page.meta_bf16_fp8)
def _run_a8w8(state: _CaseState) -> torch.Tensor:
_ensure_q_fp8(state)
return _run_aiter_decode(state, state.q_fp8, state.q_scale, state.active_page.meta_fp8_fp8)
def _resolve_dispatch(state: _CaseState) -> int:
if not _AITER_AVAILABLE:
state.dispatch_mode = 3
state.dispatch_page_size = state.active_page.page_size
return 3
if state.prefer_a8w8:
try:
_run_a8w8(state)
state.dispatch_mode = 2
state.dispatch_page_size = state.active_page.page_size
return 2
except Exception:
pass
try:
_run_a16w8(state)
state.dispatch_mode = 1
state.dispatch_page_size = state.active_page.page_size
return 1
except Exception:
if not state.a16w8_use_unit_scale:
try:
state.a16w8_use_unit_scale = True
_run_a16w8(state)
state.dispatch_mode = 1
state.dispatch_page_size = state.active_page.page_size
return 1
except Exception:
pass
else:
try:
_run_a16w8(state)
state.dispatch_mode = 1
state.dispatch_page_size = state.active_page.page_size
return 1
except Exception:
if not state.a16w8_use_unit_scale:
try:
state.a16w8_use_unit_scale = True
_run_a16w8(state)
state.dispatch_mode = 1
state.dispatch_page_size = state.active_page.page_size
return 1
except Exception:
pass
try:
_run_a8w8(state)
state.dispatch_mode = 2
state.dispatch_page_size = state.active_page.page_size
return 2
except Exception:
pass
state.dispatch_mode = 3
state.dispatch_page_size = state.active_page.page_size
return 3
def _unwrap_scaled_mm_result(result):
if isinstance(result, tuple):
return result[0]
return result
def _scaled_mm_scores(
q_fp8_2d: torch.Tensor,
k_fp8_t_2d: torch.Tensor,
q_scale: torch.Tensor,
k_scale: torch.Tensor,
) -> torch.Tensor:
scaled_mm = getattr(torch, "_scaled_mm", None)
if scaled_mm is None:
raise RuntimeError("torch._scaled_mm is unavailable")
kwargs_candidates = (
{
"scale_a": q_scale,
"scale_b": k_scale,
"out_dtype": torch.float32,
"use_fast_accum": True,
},
{
"scale_a": q_scale,
"scale_b": k_scale,
"out_dtype": torch.float32,
},
{
"scale_a": q_scale,
"scale_b": k_scale,
},
{
"scale_a": q_scale.item(),
"scale_b": k_scale.item(),
"out_dtype": torch.float32,
"use_fast_accum": True,
},
{
"scale_a": q_scale.item(),
"scale_b": k_scale.item(),
"out_dtype": torch.float32,
},
{
"scale_a": q_scale.item(),
"scale_b": k_scale.item(),
},
)
last_error = None
for kwargs in kwargs_candidates:
try:
return _unwrap_scaled_mm_result(scaled_mm(q_fp8_2d, k_fp8_t_2d, **kwargs))
except Exception as exc: # pragma: no cover - runtime compatibility path
last_error = exc
raise RuntimeError(f"torch._scaled_mm call failed: {last_error}")
def _segment_attention_with_fp8_scores(
q_bf16: torch.Tensor,
q_fp8: torch.Tensor,
q_scale: torch.Tensor,
kv_fp8: torch.Tensor,
kv_scale: torch.Tensor,
) -> torch.Tensor:
seg_q, num_heads, _ = q_bf16.shape
kv_len = kv_fp8.shape[0]
if kv_len == 0:
return torch.zeros((seg_q, num_heads, V_HEAD_DIM), dtype=torch.bfloat16, device=q_bf16.device)
q_fp8_2d = q_fp8.reshape(seg_q * num_heads, QK_HEAD_DIM)
if not q_fp8_2d.is_contiguous():
q_fp8_2d = q_fp8_2d.contiguous()
k_fp8_t_2d = kv_fp8[:, 0, :QK_HEAD_DIM].transpose(0, 1)
if not k_fp8_t_2d.is_contiguous():
k_fp8_t_2d = k_fp8_t_2d.contiguous()
scores = _scaled_mm_scores(q_fp8_2d, k_fp8_t_2d, q_scale, kv_scale)
scores = scores.to(torch.float32).mul_(SM_SCALE)
weights = torch.softmax(scores, dim=-1)
v_f32 = kv_fp8[:, 0, :V_HEAD_DIM].to(torch.float32).mul_(kv_scale)
out = weights @ v_f32
return out.view(seg_q, num_heads, V_HEAD_DIM)
def _segment_attention_fallback(
q_bf16: torch.Tensor,
kv_fp8: torch.Tensor,
kv_scale: torch.Tensor,
) -> torch.Tensor:
seg_q, num_heads, _ = q_bf16.shape
kv_len = kv_fp8.shape[0]
if kv_len == 0:
return torch.zeros((seg_q, num_heads, V_HEAD_DIM), dtype=torch.bfloat16, device=q_bf16.device)
k_f32 = kv_fp8[:, 0, :QK_HEAD_DIM].to(torch.float32).mul_(kv_scale)
v_f32 = kv_fp8[:, 0, :V_HEAD_DIM].to(torch.float32).mul_(kv_scale)
scores = torch.matmul(q_bf16.to(torch.float32), k_f32.transpose(0, 1)).mul_(SM_SCALE)
weights = torch.softmax(scores, dim=-1)
out = torch.matmul(weights, v_f32)
return out.view(seg_q, num_heads, V_HEAD_DIM)
def _run_torch_fallback(state: _CaseState) -> torch.Tensor:
_ensure_q_fp8(state)
out = state.out
batch_size = state.batch_size
q_seq_len = state.q_seq_len
kv_seq_len = state.kv_seq_len
q_bf16 = state.q_bf16
q_fp8 = state.q_fp8
q_scale = state.q_scale
kv_fp8 = state.kv_fp8
kv_scale = state.kv_scale
use_fp8_scores = True
for batch_idx in range(batch_size):
q_start = batch_idx * q_seq_len
q_end = q_start + q_seq_len
kv_start = batch_idx * kv_seq_len
kv_end = kv_start + kv_seq_len
q_seg_bf16 = q_bf16[q_start:q_end]
kv_seg_fp8 = kv_fp8[kv_start:kv_end]
if use_fp8_scores:
try:
out[q_start:q_end] = _segment_attention_with_fp8_scores(
q_seg_bf16,
q_fp8[q_start:q_end],
q_scale,
kv_seg_fp8,
kv_scale,
)
continue
except Exception: # pragma: no cover - runtime fallback path
use_fp8_scores = False
out[q_start:q_end] = _segment_attention_fallback(q_seg_bf16, kv_seg_fp8, kv_scale)
return out
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, _kv_indptr, _config = data
kv_fp8, kv_scale = kv_data["fp8"]
batch_size = qo_indptr.numel() - 1
q_seq_len = q.shape[0] // batch_size
kv_seq_len = kv_fp8.shape[0] // batch_size
state = _get_case_state(batch_size, q_seq_len, kv_seq_len, q.device)
_bind_inputs(state, q, kv_fp8, kv_scale)
mode = state.dispatch_mode
if state.dispatch_page_size != state.active_page.page_size:
mode = 0
if mode == 0:
mode = _resolve_dispatch(state)
if mode == 1:
return _run_a16w8(state)
if mode == 2:
return _run_a8w8(state)
return _run_torch_fallback(state)scrolls · 687 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