submission 590376
Jońs · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 239 lines, June 9 Researcher Reciprocity License v1.0.
Submission_v1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-590376?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:79e9449379237b9d22626a2c713d1a2d5b40eb9d85ba79899560f12e92c1e1d7
license declaredunknown
license concludedunknown
authorsJońs
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MLA decode candidate: chunked-batch MXFP4 path for q_seq_len=1.Kernel source
Submission_v1.py239 lines
#!POPCORN leaderboard amd-mixed-mla
"""
MLA decode candidate: chunked-batch MXFP4 path for q_seq_len=1.
This keeps the existing fallback paths, but replaces the slow per-segment Python
loop in the MXFP4 path with a chunked batched decode path that exploits the
uniform decode layout used by the benchmark/task cases.
"""
import os
import torch
import torch.nn.functional as F
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes
from aiter.utility.fp4_utils import e8m0_to_f32, mxfp4_to_f32
FP8_DTYPE = aiter_dtypes.fp8
QKV_DTYPE = os.environ.get("MIXED_MLA_QKV_DTYPE", "fp8")
def custom_kernel(data: input_t) -> output_t:
if QKV_DTYPE == "fp8":
return custom_kernel_fp8(data)
if QKV_DTYPE == "mxfp4":
return custom_kernel_mxfp4(data)
if QKV_DTYPE == "bf16":
return custom_kernel_bf16(data)
raise ValueError(f"Invalid QKV_DTYPE: {QKV_DTYPE}")
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 dequantize_mxfp4_segment(
kv_fp4_segment: torch.Tensor,
kv_scale: torch.Tensor,
kv_lora_rank: int,
qk_head_dim: int,
) -> tuple[torch.Tensor, torch.Tensor]:
seg_rows = kv_fp4_segment.shape[0]
num_blocks = qk_head_dim // 32
kv_fp4_2d = kv_fp4_segment.reshape(seg_rows, qk_head_dim // 2)
kv_f32 = mxfp4_to_f32(kv_fp4_2d)
scale_f32 = e8m0_to_f32(kv_scale)
scale_f32 = scale_f32[:seg_rows, :num_blocks]
scale_f32 = scale_f32.repeat_interleave(32, dim=-1)[:, :qk_head_dim]
kv_dequant = kv_f32 * scale_f32
return kv_dequant, kv_dequant[:, :kv_lora_rank]
def _uniform_decode_layout(
q: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
config: dict,
) -> tuple[int, int] | None:
batch_size = int(config["batch_size"])
q_seq_len = int(config["q_seq_len"])
kv_seq_len = int(config["kv_seq_len"])
if q_seq_len != 1:
return None
if q.shape[0] != batch_size:
return None
if qo_indptr.numel() != batch_size + 1 or kv_indptr.numel() != batch_size + 1:
return None
if int(qo_indptr[0].item()) != 0 or int(kv_indptr[0].item()) != 0:
return None
if not bool(torch.all((qo_indptr[1:] - qo_indptr[:-1]) == q_seq_len).item()):
return None
if not bool(torch.all((kv_indptr[1:] - kv_indptr[:-1]) == kv_seq_len).item()):
return None
return batch_size, kv_seq_len
def _batch_chunk_size(batch_size: int, kv_seq_len: int) -> int:
if kv_seq_len >= 8192:
return min(batch_size, 4)
if kv_seq_len >= 4096:
return min(batch_size, 8)
return min(batch_size, 32)
def _custom_kernel_mxfp4_uniform(data: input_t, batch_size: int, kv_seq_len: int) -> output_t:
q, kv_data, _qo_indptr, _kv_indptr, config = data
kv_lora_rank = config["kv_lora_rank"]
qk_head_dim = config["qk_head_dim"]
num_heads = config["num_heads"]
sm_scale = config["sm_scale"]
kv_buffer_mxfp4, kv_scale_mxfp4 = kv_data["mxfp4"]
q_batched = q.view(batch_size, num_heads, qk_head_dim).float()
kv_fp4_batched = kv_buffer_mxfp4.view(batch_size, kv_seq_len, qk_head_dim // 2)
kv_scale_batched = kv_scale_mxfp4.view(batch_size, kv_seq_len, -1)
num_blocks = qk_head_dim // 32
chunk_size = _batch_chunk_size(batch_size, kv_seq_len)
out_chunks: list[torch.Tensor] = []
for start in range(0, batch_size, chunk_size):
end = min(start + chunk_size, batch_size)
chunk_q = q_batched[start:end]
chunk_fp4 = kv_fp4_batched[start:end].reshape(-1, qk_head_dim // 2)
chunk_scale = kv_scale_batched[start:end].reshape(-1, kv_scale_batched.shape[-1])
kv_f32 = mxfp4_to_f32(chunk_fp4).view(end - start, kv_seq_len, qk_head_dim)
scale_f32 = e8m0_to_f32(chunk_scale)
scale_f32 = scale_f32[:, :num_blocks].view(end - start, kv_seq_len, num_blocks)
scale_f32 = scale_f32.repeat_interleave(32, dim=-1)[..., :qk_head_dim]
k = kv_f32 * scale_f32
v = k[..., :kv_lora_rank]
scores = torch.einsum("bhd,bsd->bhs", chunk_q * sm_scale, k)
scores = F.softmax(scores, dim=-1)
out = torch.einsum("bhs,bsv->bhv", scores, v)
out_chunks.append(out.to(torch.bfloat16))
return torch.cat(out_chunks, dim=0)
def custom_kernel_bf16(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
kv_lora_rank = config["kv_lora_rank"]
sm_scale = config["sm_scale"]
kv_buffer_bf16 = kv_data["bf16"]
batch_size = qo_indptr.shape[0] - 1
out_list = []
for i in range(batch_size):
q_s, q_e = int(qo_indptr[i].item()), int(qo_indptr[i + 1].item())
kv_s, kv_e = int(kv_indptr[i].item()), int(kv_indptr[i + 1].item())
qi = q[q_s:q_e]
kvc = kv_buffer_bf16[kv_s:kv_e, 0]
ki = kvc
vi = kvc[:, :kv_lora_rank]
qi_t = qi.float().permute(1, 0, 2)
scores = torch.matmul(qi_t * sm_scale, ki.float().T)
scores = F.softmax(scores, dim=-1)
oi = torch.matmul(scores, vi.float())
out_list.append(oi.permute(1, 0, 2).to(torch.bfloat16))
return torch.cat(out_list, dim=0)
def custom_kernel_fp8(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
num_heads = config["num_heads"]
kv_lora_rank = config["kv_lora_rank"]
qk_head_dim = config["qk_head_dim"]
sm_scale = config["sm_scale"]
kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
kv_fp8_2d = kv_buffer_fp8.view(-1, qk_head_dim)
q_fp8, q_scale = quantize_fp8(q)
batch_size = qo_indptr.shape[0] - 1
out_list = []
for i in range(batch_size):
q_s, q_e = int(qo_indptr[i].item()), int(qo_indptr[i + 1].item())
kv_s, kv_e = int(kv_indptr[i].item()), 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)
ki_fp8 = kv_fp8_2d[kv_s:kv_e]
raw_scores = torch._scaled_mm(
qi_fp8,
ki_fp8.t(),
scale_a=q_scale,
scale_b=kv_scale_fp8,
out_dtype=torch.float32,
)
scores = raw_scores.view(seq_q, num_heads, seq_kv).permute(1, 0, 2)
scores = F.softmax(scores * sm_scale, dim=-1)
vi = kv_data["bf16"][kv_s:kv_e, 0, :kv_lora_rank].float()
oi = torch.matmul(scores, vi)
out_list.append(oi.permute(1, 0, 2).to(torch.bfloat16))
return torch.cat(out_list, dim=0)
def custom_kernel_mxfp4(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
layout = _uniform_decode_layout(q, qo_indptr, kv_indptr, config)
if layout is not None:
batch_size, kv_seq_len = layout
return _custom_kernel_mxfp4_uniform(data, batch_size, kv_seq_len)
kv_lora_rank = config["kv_lora_rank"]
qk_head_dim = config["qk_head_dim"]
sm_scale = config["sm_scale"]
kv_buffer_mxfp4, kv_scale_mxfp4 = kv_data["mxfp4"]
batch_size = qo_indptr.shape[0] - 1
out_list = []
for i in range(batch_size):
q_s, q_e = int(qo_indptr[i].item()), int(qo_indptr[i + 1].item())
kv_s, kv_e = int(kv_indptr[i].item()), int(kv_indptr[i + 1].item())
qi = q[q_s:q_e]
kv_seg_fp4 = kv_buffer_mxfp4[kv_s:kv_e]
kv_seg_scale = kv_scale_mxfp4[kv_s:kv_e]
ki, vi = dequantize_mxfp4_segment(
kv_seg_fp4,
kv_seg_scale,
kv_lora_rank=kv_lora_rank,
qk_head_dim=qk_head_dim,
)
qi_t = qi.float().permute(1, 0, 2)
scores = torch.matmul(qi_t * sm_scale, ki.T)
scores = F.softmax(scores, dim=-1)
oi = torch.matmul(scores, vi)
out_list.append(oi.permute(1, 0, 2).to(torch.bfloat16))
return torch.cat(out_list, dim=0)
scrolls · 239 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