Skip to content
KernelIndex
Search⌘K

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
AMD Instinct MI355X
4.66ms
#754 of 766
2026-03-19

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.

fp4MLA 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