submission 586648
Ronit Kapoor · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 329 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-586648?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:ba5b8c6d84eeae0d7d81e82d946809dcd0e338e06482219fa398b4eb6249b487
license declaredunknown
license concludedunknown
authorsRonit Kapoor
imported2026-08-26
Kernel source
submission.py329 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
MLA (Multi-head Latent Attention) decode kernel.
Next experiment:
- bf16 default
- first try ROCm SDPA for the uniform decode case:
* q_seq_len == 1 for every request
* kv_seq_len is uniform across the batch
* dense packed q / kv buffers
- if SDPA is unavailable or slower-path conditions are not met, fall back to:
* batched matmul fast path for short KV (<= 1024)
* segmented fallback for long KV
- fp8 path kept unchanged
"""
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
QKV_DTYPE = "bf16"
FASTPATH_MAX_KV = 1024
# ---------------------------------------------------------------------------
# Dispatcher
# ---------------------------------------------------------------------------
def custom_kernel(data: input_t) -> output_t:
if QKV_DTYPE == "fp8":
return custom_kernel_fp8(data)
elif QKV_DTYPE == "bf16":
return custom_kernel_bf16(data)
else:
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 _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
# ---------------------------------------------------------------------------
# Uniform decode fast path using SDPA (best next experiment)
# ---------------------------------------------------------------------------
def _try_uniform_bf16_sdpa_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:
return None
try:
# q: [B, 16, 1, 576]
q_sdpa = q.contiguous().unsqueeze(2)
# kv: [B, L, 576]
kv_bld = kv_buffer_bf16[:, 0, :].contiguous().view(batch_size, kv_len, q.shape[-1])
# k: [B, 1, L, 576], v: [B, 1, L, 512]
k_sdpa = kv_bld.unsqueeze(1)
v_sdpa = kv_bld[:, :, :v_head_dim].unsqueeze(1)
out = F.scaled_dot_product_attention(
q_sdpa,
k_sdpa,
v_sdpa,
attn_mask=None,
dropout_p=0.0,
is_causal=False,
scale=sm_scale,
enable_gqa=True,
) # [B, 16, 1, 512]
return out.squeeze(2).to(torch.bfloat16).contiguous()
except Exception:
return None
# ---------------------------------------------------------------------------
# Existing short-KV batched fast path
# ---------------------------------------------------------------------------
def _try_uniform_bf16_matmul_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()
# ---------------------------------------------------------------------------
# Segmented bf16 fallback
# ---------------------------------------------------------------------------
def _segmented_bf16_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())
qi = q[q_s:q_e] # [seq_q, 16, 576]
kvc = kv_buffer_bf16[kv_s:kv_e, 0] # [seq_kv, 576]
ki = kvc
vi = kvc[:, :v_head_dim]
qi_t = qi.float().permute(1, 0, 2) # [16, seq_q, 576]
scores = torch.matmul(qi_t * sm_scale, ki.float().T)
probs = F.softmax(scores, dim=-1)
oi = torch.matmul(probs, vi.float()) # [16, seq_q, 512]
oi = oi.permute(1, 0, 2) # [seq_q, 16, 512]
out_list.append(oi.to(torch.bfloat16))
return torch.cat(out_list, dim=0)
# ---------------------------------------------------------------------------
# Public bf16 kernel
# ---------------------------------------------------------------------------
def custom_kernel_bf16(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"]
# 1) Best next experiment: ROCm SDPA with GQA.
fast = _try_uniform_bf16_sdpa_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
# 2) Proven short-KV fast path.
fast = _try_uniform_bf16_matmul_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
# 3) Generic fallback.
return _segmented_bf16_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 path
# ---------------------------------------------------------------------------
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 · 329 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