Skip to content
KernelIndex
Search⌘K

submission 694789

DiegoCao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 207 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-694789?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
63.3µs
#270 of 766
2026-04-02

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:0a2a04ce0ef704e140f6184e02948c120b9455c741db5d6697f15b0e5500a8d5
license declaredunknown
license concludedunknown
authorsDiegoCao
imported2026-08-26

Kernel source

submission.py207 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""
MLA Decode v10: Single-launch FP8 quant for small tensors.

For small batch sizes (bs=4, Q has only 36K elements), the ENTIRE Q tensor
fits in a single Triton thread block. This enables fusing amax reduction +
quantization into 1 launch instead of 3 (zero + amax + quant).

For large tensors: fall back to 2-launch approach (atomic amax + quant).
Direct ASM kernel calls with pre-allocated buffers for attention.
"""

import torch
import triton
import triton.language as tl
from task import input_t, output_t

import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

_FP8 = aiter_dtypes.fp8
_FP8_MAX = float(torch.finfo(_FP8).max)
_SM = 1.0 / (576.0 ** 0.5)


# ---------------------------------------------------------------------------
# Single-block fused amax+quant kernel (1 launch for small tensors)
# ---------------------------------------------------------------------------
@triton.jit
def _fused_quant_small_kernel(
    X_ptr, Out_ptr, Scale_ptr,
    N: tl.constexpr, FP8_MAX: tl.constexpr, BLOCK: tl.constexpr,
):
    """Single-block: compute amax + quantize in one pass."""
    off = tl.arange(0, BLOCK)
    mask = off < N
    x = tl.load(X_ptr + off, mask=mask, other=0.0).to(tl.float32)

    # Global amax across entire block
    amax = tl.max(tl.abs(x))
    amax = tl.maximum(amax, 1e-12)
    inv_scale = FP8_MAX / amax
    tl.store(Scale_ptr, amax / FP8_MAX)

    # Quantize
    q = x * inv_scale
    q = tl.minimum(tl.maximum(q, -FP8_MAX), FP8_MAX)
    tl.store(Out_ptr + off, q.to(Out_ptr.dtype.element_ty), mask=mask)


# ---------------------------------------------------------------------------
# Multi-block atomic amax + quant (2 launches for large tensors)
# ---------------------------------------------------------------------------
_BLOCK_LARGE = 4096

@triton.jit
def _amax_kernel(X_ptr, Amax_ptr, N: tl.constexpr, BLOCK: tl.constexpr):
    pid = tl.program_id(0)
    off = pid * BLOCK + tl.arange(0, BLOCK)
    mask = off < N
    x = tl.load(X_ptr + off, mask=mask, other=0.0).to(tl.float32)
    local_max = tl.max(tl.abs(x))
    tl.atomic_max(Amax_ptr, local_max)

@triton.jit
def _quant_kernel(
    X_ptr, Out_ptr, Amax_ptr, Scale_ptr, N: tl.constexpr,
    FP8_MAX: tl.constexpr, BLOCK: tl.constexpr,
):
    pid = tl.program_id(0)
    off = pid * BLOCK + tl.arange(0, BLOCK)
    mask = off < N
    amax = tl.load(Amax_ptr).to(tl.float32)
    amax = tl.maximum(amax, 1e-12)
    inv_scale = FP8_MAX / amax
    if pid == 0:
        tl.store(Scale_ptr, amax / FP8_MAX)
    x = tl.load(X_ptr + off, mask=mask, other=0.0).to(tl.float32)
    q = x * inv_scale
    q = tl.minimum(tl.maximum(q, -FP8_MAX), FP8_MAX)
    tl.store(Out_ptr + off, q.to(Out_ptr.dtype.element_ty), mask=mask)


# ---------------------------------------------------------------------------
# Quantization dispatch
# ---------------------------------------------------------------------------
# Threshold for single-block kernel: N must fit in one block
# Max block size ~65536 elements (16 warps, safe for AMD CDNA4)
_SINGLE_BLOCK_MAX = 65536

_quant_bufs: dict = {}


def _quant_q(q):
    """FP8 quantize Q: 1 launch for small tensors, 3 launches for large."""
    N = q.numel()
    c = _quant_bufs.get(N)
    if c is None:
        c = {
            "out": torch.empty(N, dtype=_FP8, device="cuda"),
            "amax": torch.zeros(1, dtype=torch.float32, device="cuda"),
            "scale": torch.empty(1, dtype=torch.float32, device="cuda"),
        }
        _quant_bufs[N] = c

    if N <= _SINGLE_BLOCK_MAX:
        # Single-block fused kernel: 1 launch!
        BLOCK = triton.next_power_of_2(N)
        _fused_quant_small_kernel[(1,)](
            q, c["out"], c["scale"],
            N=N, FP8_MAX=_FP8_MAX, BLOCK=BLOCK,
            num_warps=min(16, max(1, BLOCK // 256)),
        )
    else:
        # Multi-block: 3 launches (zero + amax + quant)
        c["amax"].zero_()
        grid = ((N + _BLOCK_LARGE - 1) // _BLOCK_LARGE,)
        _amax_kernel[grid](q, c["amax"], N=N, BLOCK=_BLOCK_LARGE)
        _quant_kernel[grid](q, c["out"], c["amax"], c["scale"], N=N, FP8_MAX=_FP8_MAX, BLOCK=_BLOCK_LARGE)

    return c["out"].view(q.shape), c["scale"]


# ---------------------------------------------------------------------------
# Attention state cache
# ---------------------------------------------------------------------------
_SPLITS = {
    (4, 1024): 8, (4, 8192): 16,
    (32, 1024): 16, (32, 8192): 32,
    (64, 1024): 16, (64, 8192): 32,
    (256, 1024): 16, (256, 8192): 32,
}

_cache: dict = {}


class _State:
    __slots__ = (
        'nks', 'kv_indices', 'kv_lpl', 'output',
        'wm', 'wi', 'wis', 'ri', 'rfm', 'rpm',
        'logits', 'attn_lse', 'final_lse',
    )

    def __init__(self, bs, kvlen, q_dtype, kv_dtype, qo_indptr, kv_indptr):
        total_kv = bs * kvlen
        self.nks = _SPLITS.get((bs, kvlen), 32)
        self.kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
        self.kv_lpl = torch.full((bs,), kvlen, dtype=torch.int32, device="cuda")
        self.output = torch.empty((bs, 16, 512), dtype=torch.bfloat16, device="cuda")

        info = get_mla_metadata_info_v1(
            bs, 1, 16, q_dtype, kv_dtype,
            is_sparse=False, fast_mode=False,
            num_kv_splits=self.nks, intra_batch_mode=True,
        )
        work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
        self.wm, self.wi, self.wis, self.ri, self.rfm, self.rpm = work
        get_mla_metadata_v1(
            qo_indptr, kv_indptr, self.kv_lpl,
            16, 1, True,
            self.wm, self.wis, self.wi, self.ri, self.rfm, self.rpm,
            page_size=1, kv_granularity=16,
            max_seqlen_qo=1, uni_seqlen_qo=1,
            fast_mode=False, max_split_per_batch=self.nks,
            intra_batch_mode=True, dtype_q=q_dtype, dtype_kv=kv_dtype,
        )
        rpm_size = self.rpm.size(0)
        self.logits = torch.empty((rpm_size, 1, 16, 512), dtype=torch.float32, device="cuda")
        self.attn_lse = torch.empty((rpm_size, 1, 16, 1), dtype=torch.float32, device="cuda")
        self.final_lse = torch.empty((bs, 16), dtype=torch.float32, device="cuda")


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = config["batch_size"]
    kvlen = config["kv_seq_len"]
    kv_buf, kv_sc = kv_data["fp8"]

    q_fp8, q_sc = _quant_q(q)

    key = (bs, kvlen)
    s = _cache.get(key)
    if s is None:
        s = _State(bs, kvlen, q_fp8.dtype, kv_buf.dtype, qo_indptr, kv_indptr)
        _cache[key] = s

    # Direct ASM kernel calls with pre-allocated buffers
    aiter.mla_decode_stage1_asm_fwd(
        q_fp8.view(-1, 16, 576),
        kv_buf.view(-1, 1, 1, 576),
        qo_indptr, kv_indptr,
        s.kv_indices, s.kv_lpl,
        None, s.wm, s.wi, s.wis,
        1, 1, 1, _SM,
        s.logits, s.attn_lse, s.output,
        q_sc, kv_sc,
    )
    aiter.mla_reduce_v1(
        s.logits, s.attn_lse,
        s.ri, s.rfm, s.rpm,
        1, s.output, s.final_lse,
    )
    return s.output
scrolls · 207 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