Skip to content
KernelIndex
Search⌘K

submission 667881

migratesky · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v153_prealloc_tuned.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-667881?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
65.0µs
#284 of 766
2026-03-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:7bcce6a066905200c24a567a1f542f6cc702c690e673056131b1ac82e1ef5d97
license declaredunknown
license concludedunknown
authorsmigratesky
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

mmascores += tl.dot(q_chunk, tl.trans(kv_chunk.to(tl.bfloat16)))
num-warps = 4num_warps=4,
online-softmaxm_new = tl.maximum(m_i, m_ij)
stages = 2num_stages = 2 if kv_seq_len < 4096 else 1
tile-n = 64BLOCK_N = 64 if kv_seq_len >= 4096 else 32

Kernel source

v153_prealloc_tuned.py346 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""v153: Pre-allocated buffers + per-shape warp tuning.

Base: v146 (geo mean 68.019us).

Changes vs v146:
1. Pre-allocate partial_o, partial_lse, out buffers once (save ~3-5us
   per call from avoiding torch.empty).
2. B<=32: use 4 warps (v121 config - marginally better for B32/KV8192).
3. B>=64: keep 8 warps (v146 config - proven best for large batches).
4. BLOCK_V=512 for reduce kernel (single V iteration vs 4 with BLOCK_V=128).
"""

from __future__ import annotations

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

_LOGGED = False
_PARTIAL_O = None
_PARTIAL_LSE = None
_OUT_BUF = None


def _log(msg: str):
    print(f"[v153] {msg}", file=sys.stderr, flush=True)


def _ensure_buffers(device):
    global _PARTIAL_O, _PARTIAL_LSE, _OUT_BUF
    if _PARTIAL_O is None:
        _PARTIAL_O = torch.empty(
            (2048, 16, 512), dtype=torch.float32, device=device,
        )
        _PARTIAL_LSE = torch.empty(
            (2048, 16), dtype=torch.float32, device=device,
        )
        _OUT_BUF = torch.empty(
            (256, 16, 512), dtype=torch.bfloat16, device=device,
        )


@triton.jit
def _flash_decode_stage1(
    Q_ptr, KV_ptr, KV_scale_ptr,
    Partial_O_ptr, Partial_LSE_ptr,
    qo_indptr_ptr, kv_indptr_ptr,
    sm_scale,
    stride_q_tok, stride_q_h, stride_q_d,
    stride_kv_tok, stride_kv_d,
    stride_po_row, stride_po_h, stride_po_d,
    stride_plse_row, stride_plse_h,
    NUM_KV_SPLITS: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_D: tl.constexpr,
    NUM_HEADS: tl.constexpr,
    QK_DIM: tl.constexpr,
    V_DIM: tl.constexpr,
):
    batch_id = tl.program_id(0)
    split_id = tl.program_id(1)

    kv_start = tl.load(kv_indptr_ptr + batch_id)
    kv_end = tl.load(kv_indptr_ptr + batch_id + 1)
    kv_len = kv_end - kv_start

    tokens_per_split = tl.cdiv(kv_len, NUM_KV_SPLITS)
    split_start = split_id * tokens_per_split
    split_end = tl.minimum(split_start + tokens_per_split, kv_len)

    row = batch_id * NUM_KV_SPLITS + split_id

    if split_start >= kv_len:
        offs_h = tl.arange(0, NUM_HEADS)
        tl.store(
            Partial_LSE_ptr + row * stride_plse_row
            + offs_h * stride_plse_h,
            tl.full([NUM_HEADS], float("-inf"), dtype=tl.float32),
        )
        return

    qo_start = tl.load(qo_indptr_ptr + batch_id)
    q_base = Q_ptr + qo_start * stride_q_tok
    kv_scale = tl.load(KV_scale_ptr)

    LOG2E: tl.constexpr = 1.4426950408889634
    qk_scale = sm_scale * kv_scale * LOG2E

    m_i = tl.full([NUM_HEADS], float("-inf"), dtype=tl.float32)
    l_i = tl.zeros([NUM_HEADS], dtype=tl.float32)
    acc = tl.zeros([NUM_HEADS, V_DIM], dtype=tl.float32)

    offs_h = tl.arange(0, NUM_HEADS)
    offs_n = tl.arange(0, BLOCK_N)

    for tok_start in range(split_start, split_end, BLOCK_N):
        n_valid = tl.minimum(BLOCK_N, split_end - tok_start)
        kv_tok_ids = kv_start + tok_start + offs_n
        mask_n = offs_n < n_valid

        scores = tl.zeros([NUM_HEADS, BLOCK_N], dtype=tl.float32)
        for d_start in range(0, QK_DIM, BLOCK_D):
            offs_d = d_start + tl.arange(0, BLOCK_D)
            q_chunk = tl.load(
                q_base
                + offs_h[:, None] * stride_q_h
                + offs_d[None, :] * stride_q_d,
            ).to(tl.bfloat16)

            kv_chunk = tl.load(
                KV_ptr
                + kv_tok_ids[:, None] * stride_kv_tok
                + offs_d[None, :] * stride_kv_d,
                mask=mask_n[:, None],
                other=0.0,
            )
            scores += tl.dot(q_chunk, tl.trans(kv_chunk.to(tl.bfloat16)))

        scores = scores * qk_scale
        scores = tl.where(mask_n[None, :], scores, float("-inf"))

        m_ij = tl.max(scores, axis=1)
        m_new = tl.maximum(m_i, m_ij)
        alpha = tl.math.exp2(m_i - m_new)
        p = tl.math.exp2(scores - m_new[:, None])
        l_new = alpha * l_i + tl.sum(p, axis=1)
        acc = acc * alpha[:, None]

        offs_v = tl.arange(0, V_DIM)
        v_block = tl.load(
            KV_ptr
            + kv_tok_ids[:, None] * stride_kv_tok
            + offs_v[None, :] * stride_kv_d,
            mask=mask_n[:, None],
            other=0.0,
        )
        p_bf16 = p.to(tl.bfloat16)
        v_bf16 = v_block.to(tl.bfloat16)
        acc += tl.dot(p_bf16, v_bf16)

        m_i = m_new
        l_i = l_new

    acc = (acc * kv_scale) / l_i[:, None]
    offs_v = tl.arange(0, V_DIM)
    tl.store(
        Partial_O_ptr + row * stride_po_row
        + offs_h[:, None] * stride_po_h
        + offs_v[None, :] * stride_po_d,
        acc.to(tl.float32),
    )
    lse = tl.math.log2(l_i) + m_i
    tl.store(
        Partial_LSE_ptr + row * stride_plse_row + offs_h * stride_plse_h,
        lse,
    )


@triton.jit
def _flash_decode_reduce(
    Partial_O_ptr, Partial_LSE_ptr, Out_ptr,
    stride_po_row, stride_po_h, stride_po_d,
    stride_plse_row, stride_plse_h,
    stride_o_tok, stride_o_h, stride_o_d,
    qo_indptr_ptr,
    NUM_KV_SPLITS: tl.constexpr,
    NUM_HEADS: tl.constexpr,
    V_DIM: tl.constexpr,
    BLOCK_V: tl.constexpr,
):
    batch_id = tl.program_id(0)
    head_id = tl.program_id(1)

    offs_s = tl.arange(0, NUM_KV_SPLITS)
    base_row = batch_id * NUM_KV_SPLITS
    lse_vals = tl.load(
        Partial_LSE_ptr
        + (base_row + offs_s) * stride_plse_row
        + head_id * stride_plse_h,
    )

    max_lse = tl.max(lse_vals, axis=0)
    weights = tl.math.exp2(lse_vals - max_lse)
    sum_weights = tl.sum(weights, axis=0)
    sum_weights = tl.where(sum_weights > 0.0, sum_weights, 1.0)

    qo_start = tl.load(qo_indptr_ptr + batch_id)
    offs_v = tl.arange(0, BLOCK_V)
    for v_start in range(0, V_DIM, BLOCK_V):
        v_offs = v_start + offs_v
        mask_v = v_offs < V_DIM

        partial_all = tl.load(
            Partial_O_ptr
            + (base_row + offs_s)[:, None] * stride_po_row
            + head_id * stride_po_h
            + v_offs[None, :] * stride_po_d,
            mask=mask_v[None, :], other=0.0,
        )
        weighted = partial_all * weights[:, None]
        acc = tl.sum(weighted, axis=0) / sum_weights
        tl.store(
            Out_ptr + qo_start * stride_o_tok
            + head_id * stride_o_h
            + v_offs * stride_o_d,
            acc.to(tl.bfloat16), mask=mask_v,
        )


def _choose_splits_v121(batch_size: int, kv_seq_len: int) -> int:
    CU_COUNT = 256
    BLOCK_N_EST = 64 if kv_seq_len >= 4096 else 32
    min_tokens_per_split = BLOCK_N_EST * 2
    max_splits = max(1, kv_seq_len // min_tokens_per_split)
    if batch_size <= 4:
        target_programs = CU_COUNT * 16
    elif batch_size <= 32:
        target_programs = CU_COUNT * 4
    elif batch_size <= 64:
        target_programs = CU_COUNT * 4
    else:
        target_programs = CU_COUNT * 8
    if batch_size >= target_programs:
        return 1
    ns = max(1, target_programs // batch_size)
    ns = min(ns, max_splits)
    ns = 1 << (ns - 1).bit_length() if ns > 1 else 1
    return min(ns, 128)


def _choose_splits_v146(batch_size: int, kv_seq_len: int) -> int:
    CU_COUNT = 256
    BLOCK_N_EST = 64 if kv_seq_len >= 4096 else 32
    min_tokens_per_split = BLOCK_N_EST * 2
    max_splits = max(1, kv_seq_len // min_tokens_per_split)
    if batch_size <= 32:
        target_programs = CU_COUNT * 4
    elif batch_size <= 64:
        target_programs = CU_COUNT * 4
    elif batch_size >= 128:
        target_programs = CU_COUNT * 4
    else:
        target_programs = CU_COUNT * 8
    if batch_size >= target_programs:
        return 1
    ns = max(1, target_programs // batch_size)
    ns = min(ns, max_splits)
    ns = 1 << (ns - 1).bit_length() if ns > 1 else 1
    return min(ns, 128)


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    global _LOGGED

    q, kv_data, qo_indptr, kv_indptr, config = data
    kv_fp8, kv_scale_tensor = kv_data["fp8"]

    batch_size = int(config["batch_size"])
    kv_seq_len = int(config["kv_seq_len"])
    num_heads = int(config["num_heads"])
    qk_head_dim = int(config["qk_head_dim"])
    v_head_dim = int(config["v_head_dim"])
    sm_scale = float(config["sm_scale"])
    device = q.device

    _ensure_buffers(device)

    if batch_size <= 4:
        nwarps = 4
        BLOCK_N = 64 if kv_seq_len >= 4096 else 32
        num_stages = 2 if kv_seq_len < 4096 else 1
        num_kv_splits = _choose_splits_v121(batch_size, kv_seq_len)
    elif batch_size <= 32:
        nwarps = 4
        BLOCK_N = 64 if kv_seq_len >= 4096 else 32
        num_stages = 2 if kv_seq_len < 4096 else 1
        num_kv_splits = _choose_splits_v146(batch_size, kv_seq_len)
    else:
        nwarps = 8
        BLOCK_N = 128 if kv_seq_len >= 4096 else 64
        num_stages = 2
        num_kv_splits = _choose_splits_v146(batch_size, kv_seq_len)

    if not _LOGGED:
        _log(
            f"B={batch_size} KV={kv_seq_len} splits={num_kv_splits} "
            f"warps={nwarps} BN={BLOCK_N} stages={num_stages}"
        )
        _LOGGED = True

    kv_flat = kv_fp8.view(-1, qk_head_dim)
    BLOCK_D = 64

    n_rows = batch_size * num_kv_splits
    partial_o = _PARTIAL_O[:n_rows]
    partial_lse = _PARTIAL_LSE[:n_rows]

    grid_s1 = (batch_size, num_kv_splits)
    _flash_decode_stage1[grid_s1](
        q, kv_flat, kv_scale_tensor,
        partial_o, partial_lse,
        qo_indptr, kv_indptr,
        sm_scale,
        q.stride(0), q.stride(1), q.stride(2),
        kv_flat.stride(0), kv_flat.stride(1),
        partial_o.stride(0), partial_o.stride(1), partial_o.stride(2),
        partial_lse.stride(0), partial_lse.stride(1),
        NUM_KV_SPLITS=num_kv_splits,
        BLOCK_N=BLOCK_N,
        BLOCK_D=BLOCK_D,
        NUM_HEADS=num_heads,
        QK_DIM=qk_head_dim,
        V_DIM=v_head_dim,
        num_warps=nwarps,
        num_stages=num_stages,
    )

    out = _OUT_BUF[:batch_size]

    if num_kv_splits == 1:
        out.copy_(partial_o[:batch_size].view_as(out).to(torch.bfloat16))
    else:
        BLOCK_V = 512
        grid_r = (batch_size, num_heads)
        _flash_decode_reduce[grid_r](
            partial_o, partial_lse, out,
            partial_o.stride(0), partial_o.stride(1), partial_o.stride(2),
            partial_lse.stride(0), partial_lse.stride(1),
            out.stride(0), out.stride(1), out.stride(2),
            qo_indptr,
            NUM_KV_SPLITS=num_kv_splits,
            NUM_HEADS=num_heads,
            V_DIM=v_head_dim,
            BLOCK_V=BLOCK_V,
            num_warps=4,
            num_stages=1,
        )

    return out
scrolls · 346 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