submission 586972
ron1tk · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 383 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-586972?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:f6be6ec1414e74942a7d8a5d107218a59df25823c4034fbeb3f0b9c4d7f7c62e
license declaredunknown
license concludedunknown
authorsron1tk
imported2026-08-26
Kernel source
submission.py383 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
MLA decode submission for amd-mixed-mla.
Primary strategy:
- Use AMD AITER's public mla_decode_fwd directly on dense bf16 KV by wrapping the
dense cache as page_size=1 paged KV.
- Fall back to the best-performing hybrid bf16 implementation if AITER is not
available or the call fails in the harness.
Why this is the strongest public-path attempt:
- AMD's official docs expose mla_decode_fwd as the optimized MLA decode API.
- vLLM's ROCm backend notes that the assembly mla_decode_fwd kernel is where most
decode performance gains come from.
"""
from __future__ import annotations
import torch
import torch.nn.functional as F
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes
FP8_DTYPE = aiter_dtypes.fp8
# Keep bf16 as the main data path; AITER handles the actual optimized decode.
QKV_DTYPE = "bf16"
FASTPATH_MAX_KV = 1024
try:
from aiter.mla import mla_decode_fwd as aiter_mla_decode_fwd
AITER_MLA_AVAILABLE = True
except Exception:
aiter_mla_decode_fwd = None
AITER_MLA_AVAILABLE = False
# -----------------------------------------------------------------------------
# Small caches for benchmark repeats
# -----------------------------------------------------------------------------
_KV_INDICES_CACHE: dict[tuple[int, int], torch.Tensor] = {}
_KV_LAST_PAGE_LENS_CACHE: dict[tuple[int, tuple[int, ...]], torch.Tensor] = {}
# -----------------------------------------------------------------------------
# Dispatcher
# -----------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
# First try the vendor-optimized path.
if QKV_DTYPE == "bf16":
out = _try_aiter_bf16(data)
if out is not None:
return out
return custom_kernel_bf16_hybrid(data)
if QKV_DTYPE == "fp8":
return custom_kernel_fp8(data)
raise ValueError(f"Invalid QKV_DTYPE: {QKV_DTYPE}")
# -----------------------------------------------------------------------------
# Helpers
# -----------------------------------------------------------------------------
def quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
finfo = torch.finfo(FP8_DTYPE)
amax = tensor.abs().amax().clamp(min=1e-12)
scale = amax / finfo.max
fp8_tensor = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
return fp8_tensor, scale.to(torch.float32).reshape(1)
def as_scalar_scale(x, device: torch.device) -> torch.Tensor:
if torch.is_tensor(x):
return x.to(device=device, dtype=torch.float32).reshape(1)
return torch.tensor([float(x)], device=device, dtype=torch.float32)
def _device_key(device: torch.device) -> int:
return -1 if device.index is None else int(device.index)
def _get_dense_kv_indices(total_kv: int, device: torch.device) -> torch.Tensor:
key = (_device_key(device), total_kv)
t = _KV_INDICES_CACHE.get(key)
if t is None:
t = torch.arange(total_kv, device=device, dtype=torch.int32)
_KV_INDICES_CACHE[key] = t
return t
def _get_kv_last_page_lens(kv_indptr: torch.Tensor) -> torch.Tensor:
kv_lens = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
# page_size = 1, so any non-empty sequence has last_page_len = 1
vals = torch.where(kv_lens > 0, torch.ones_like(kv_lens), torch.zeros_like(kv_lens))
key = (_device_key(kv_indptr.device), tuple(int(x) for x in vals.tolist()))
t = _KV_LAST_PAGE_LENS_CACHE.get(key)
if t is None:
t = vals.contiguous()
_KV_LAST_PAGE_LENS_CACHE[key] = t
return t
def _check_uniform_decode_layout(
q: torch.Tensor,
kv_buffer_bf16: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
) -> tuple[bool, int, int]:
batch_size = qo_indptr.numel() - 1
if batch_size <= 0:
return False, 0, 0
qo_lens = qo_indptr[1:] - qo_indptr[:-1]
kv_lens = kv_indptr[1:] - kv_indptr[:-1]
if not bool(torch.all(qo_lens == 1).item()):
return False, 0, 0
if not bool(torch.all(kv_lens == kv_lens[0]).item()):
return False, 0, 0
kv_len = int(kv_lens[0].item())
total_q = q.shape[0]
total_kv = kv_buffer_bf16.shape[0]
if total_q != batch_size:
return False, 0, 0
if total_kv != batch_size * kv_len:
return False, 0, 0
return True, batch_size, kv_len
# -----------------------------------------------------------------------------
# AITER primary path
# -----------------------------------------------------------------------------
def _try_aiter_bf16(data: input_t) -> torch.Tensor | None:
if not AITER_MLA_AVAILABLE:
return None
q, kv_data, qo_indptr, kv_indptr, config = data
# Public AITER decode currently targets the DeepSeek-style 16-head case.
if int(config["num_heads"]) not in (16, 128):
return None
kv_buffer_bf16 = kv_data["bf16"]
device = q.device
try:
total_q = q.shape[0]
qk_head_dim = q.shape[-1]
v_head_dim = int(config["v_head_dim"])
sm_scale = float(config["sm_scale"])
# Dense KV -> paged KV with page_size=1
# AITER expects [num_pages, page_size, num_heads_kv, qk_head_dim]
kv_paged = kv_buffer_bf16.contiguous().view(-1, 1, 1, qk_head_dim)
# For dense storage, the page indices are just 0..total_kv-1
total_kv = int(kv_indptr[-1].item())
kv_indices = _get_dense_kv_indices(total_kv, device)
kv_last_page_lens = _get_kv_last_page_lens(kv_indptr)
max_seqlen_q = int((qo_indptr[1:] - qo_indptr[:-1]).max().item())
if max_seqlen_q <= 0:
max_seqlen_q = 1
o = torch.empty((total_q, q.shape[1], v_head_dim), device=device, dtype=torch.bfloat16)
ret = aiter_mla_decode_fwd(
q.contiguous(),
kv_paged,
o,
qo_indptr.to(torch.int32).contiguous(),
kv_indptr.to(torch.int32).contiguous(),
kv_indices,
kv_last_page_lens,
max_seqlen_q,
sm_scale=sm_scale,
)
if torch.is_tensor(ret):
return ret.to(torch.bfloat16)
return o
except Exception:
return None
# -----------------------------------------------------------------------------
# Best fallback: short-KV batched fast path + decode-specialized long-KV fallback
# -----------------------------------------------------------------------------
def _try_uniform_bf16_fastpath(
q: torch.Tensor,
kv_buffer_bf16: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
v_head_dim: int,
sm_scale: float,
) -> torch.Tensor | None:
ok, batch_size, kv_len = _check_uniform_decode_layout(
q=q,
kv_buffer_bf16=kv_buffer_bf16,
qo_indptr=qo_indptr,
kv_indptr=kv_indptr,
)
if not ok or kv_len > FASTPATH_MAX_KV:
return None
q_bhd = q.contiguous() # [B, H, D]
kv_bld = kv_buffer_bf16[:, 0, :].contiguous().view(batch_size, kv_len, q.shape[-1])
k_bld = kv_bld
v_blv = kv_bld[:, :, :v_head_dim]
scores = torch.matmul(
(q_bhd * sm_scale).unsqueeze(2), # [B, H, 1, D]
k_bld.unsqueeze(1).transpose(-1, -2), # [B, 1, D, L]
).squeeze(2) # [B, H, L]
probs = F.softmax(scores.float(), dim=-1).to(q_bhd.dtype)
out = torch.matmul(
probs.unsqueeze(2), # [B, H, 1, L]
v_blv.unsqueeze(1), # [B, 1, L, V]
).squeeze(2) # [B, H, V]
return out.to(torch.bfloat16).contiguous()
def _segmented_bf16_decode_fallback(
q: torch.Tensor,
kv_buffer_bf16: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
v_head_dim: int,
sm_scale: float,
) -> torch.Tensor:
batch_size = qo_indptr.shape[0] - 1
out_list = []
for i in range(batch_size):
q_s = int(qo_indptr[i].item())
q_e = int(qo_indptr[i + 1].item())
kv_s = int(kv_indptr[i].item())
kv_e = int(kv_indptr[i + 1].item())
seq_q = q_e - q_s
if seq_q == 0:
continue
if seq_q == 1:
qh = q[q_s, :, :] # [H, D] bf16
kvc = kv_buffer_bf16[kv_s:kv_e, 0] # [L, D] bf16
scores = torch.matmul(qh * sm_scale, kvc.transpose(0, 1)) # [H, L] bf16
probs = F.softmax(scores.float(), dim=-1).to(qh.dtype)
vi = kvc[:, :v_head_dim] # [L, V] bf16
out_hv = torch.matmul(probs, vi) # [H, V] bf16
out_list.append(out_hv.unsqueeze(0).to(torch.bfloat16))
continue
qi = q[q_s:q_e] # [seq_q, H, D]
kvc = kv_buffer_bf16[kv_s:kv_e, 0] # [L, D]
qi_t = qi.permute(1, 0, 2).contiguous() # [H, seq_q, D]
ki_t = kvc.transpose(0, 1).contiguous() # [D, L]
scores = torch.matmul(qi_t * sm_scale, ki_t)
probs = F.softmax(scores.float(), dim=-1).to(qi.dtype)
vi = kvc[:, :v_head_dim]
oi = torch.matmul(probs, vi)
oi = oi.permute(1, 0, 2).contiguous()
out_list.append(oi.to(torch.bfloat16))
if not out_list:
return torch.empty((0, q.shape[1], v_head_dim), device=q.device, dtype=torch.bfloat16)
return torch.cat(out_list, dim=0)
def custom_kernel_bf16_hybrid(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
v_head_dim = int(config["v_head_dim"])
sm_scale = float(config["sm_scale"])
kv_buffer_bf16 = kv_data["bf16"]
fast = _try_uniform_bf16_fastpath(
q=q,
kv_buffer_bf16=kv_buffer_bf16,
qo_indptr=qo_indptr,
kv_indptr=kv_indptr,
v_head_dim=v_head_dim,
sm_scale=sm_scale,
)
if fast is not None:
return fast
return _segmented_bf16_decode_fallback(
q=q,
kv_buffer_bf16=kv_buffer_bf16,
qo_indptr=qo_indptr,
kv_indptr=kv_indptr,
v_head_dim=v_head_dim,
sm_scale=sm_scale,
)
# -----------------------------------------------------------------------------
# FP8 fallback path (kept for experimentation)
# -----------------------------------------------------------------------------
def custom_kernel_fp8(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
num_heads = int(config["num_heads"])
v_head_dim = int(config["v_head_dim"])
qk_head_dim = int(config["qk_head_dim"])
sm_scale = float(config["sm_scale"])
kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
kv_scale_fp8 = as_scalar_scale(kv_scale_fp8, q.device)
kv_fp8_2d = kv_buffer_fp8.reshape(-1, qk_head_dim)
kv_bf16 = kv_data["bf16"]
q_fp8, q_scale = quantize_fp8(q)
q_scale = as_scalar_scale(q_scale, q.device)
batch_size = qo_indptr.shape[0] - 1
out_list = []
for i in range(batch_size):
q_s = int(qo_indptr[i].item())
q_e = int(qo_indptr[i + 1].item())
kv_s = int(kv_indptr[i].item())
kv_e = int(kv_indptr[i + 1].item())
seq_q = q_e - q_s
seq_kv = kv_e - kv_s
qi_fp8 = q_fp8[q_s:q_e].reshape(seq_q * num_heads, qk_head_dim).contiguous()
ki_fp8 = kv_fp8_2d[kv_s:kv_e].contiguous()
try:
raw_scores = torch._scaled_mm(
qi_fp8,
ki_fp8,
scale_a=q_scale,
scale_b=kv_scale_fp8,
out_dtype=torch.float32,
)
except RuntimeError as e:
msg = str(e)
if "cuBLASLt" not in msg and "_scaled_mm" not in msg:
raise
raw_scores = (qi_fp8.float() * q_scale).matmul(
(ki_fp8.float() * kv_scale_fp8).transpose(0, 1)
)
scores = raw_scores.view(seq_q, num_heads, seq_kv).permute(1, 0, 2)
scores = scores * sm_scale
probs = F.softmax(scores, dim=-1)
vi = kv_bf16[kv_s:kv_e, 0, :v_head_dim].float()
oi = torch.matmul(probs, vi)
oi = oi.permute(1, 0, 2).to(torch.bfloat16)
out_list.append(oi)
return torch.cat(out_list, dim=0)scrolls · 383 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