Skip to content
KernelIndex
Search⌘K

submission 653127

mocimex265 · 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.

test_v96_leaderboard_safe.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-653127?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
55.3µs
#187 of 766
2026-03-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:40aacdaa82e74a5b808caa9422ab228350d4f22282d40c33ce13300804ea75c2
license declaredunknown
license concludedunknown
authorsmocimex265
imported2026-08-15

Techniques

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

num-warps = 4num_warps=4,
persistent-kernelTarget: s4/s6 — revert to persistent (nonpers fails leaderboard for kv=8192)

Kernel source

test_v96_leaderboard_safe.py329 lines
"""
test_v96_leaderboard_safe: Revert nonpers to kv<=1024 only (fix leaderboard s4 failure)
Base: test_v94_gran64_all.py
Direction: NEW — fix leaderboard correctness
Target: s4/s6 — revert to persistent (nonpers fails leaderboard for kv=8192)
Change: Restore kv_seq_len<=1024 condition for nonpers fp8 tier. s4/s6 back to
        persistent with kv_gran=64. Keep nonpers for s3/s5 (kv<=1024, proven safe).
        v94 leaderboard failed on s4 (batch=32, kv=8192) — nonpers + custom reduce
        breaks for kv=8192 in leaderboard mode (same pattern as v77-v78).
Rationale: v94 leaderboard failed on s4. Must revert kv>1024 to persistent.
Scale: INCREMENTAL
"""
import torch
import aiter
import triton
import triton.language as tl
from task import input_t, output_t

from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.mla import get_meta_param

NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM
V_HEAD_DIM = KV_LORA_RANK
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8

_cache = {}


# ========================================================================
# Custom Triton reduce kernel (replaces aiter's _fwd_kernel_stage2_asm)
# Simple, stateless, no caching — should be leaderboard-safe.
# ========================================================================

@triton.jit
def _custom_reduce_kernel(
    logits_ptr,   # float32 (total_q, num_splits, num_heads, v_dim)
    lse_ptr,      # float32 (total_q, num_splits, num_heads, 1)
    output_ptr,   # bf16 (total_q, num_heads, v_dim)
    total_q: tl.int32,
    num_splits: tl.constexpr,
    num_heads: tl.constexpr,
    v_dim: tl.constexpr,
    BLOCK_V: tl.constexpr,
):
    # Grid: (total_q, num_heads, cdiv(v_dim, BLOCK_V))
    q_idx = tl.program_id(0)
    head_idx = tl.program_id(1)
    v_block = tl.program_id(2)

    v_offs = v_block * BLOCK_V + tl.arange(0, BLOCK_V)
    v_mask = v_offs < v_dim

    # Strides for logits: (total_q, num_splits, num_heads, v_dim)
    logits_q_stride = num_splits * num_heads * v_dim
    logits_s_stride = num_heads * v_dim
    logits_h_stride = v_dim

    # Strides for lse: (total_q, num_splits, num_heads, 1)
    lse_q_stride = num_splits * num_heads
    lse_s_stride = num_heads

    # Find max LSE across splits for numerical stability
    max_lse = tl.full((), float("-inf"), dtype=tl.float32)
    for s in range(num_splits):
        lse_val = tl.load(lse_ptr + q_idx * lse_q_stride + s * lse_s_stride + head_idx)
        max_lse = tl.maximum(max_lse, lse_val)

    # Weighted sum of logits across splits
    acc = tl.zeros((BLOCK_V,), dtype=tl.float32)
    weight_sum = tl.full((), 0.0, dtype=tl.float32)

    for s in range(num_splits):
        lse_val = tl.load(lse_ptr + q_idx * lse_q_stride + s * lse_s_stride + head_idx)
        w = tl.exp(lse_val - max_lse)
        weight_sum += w

        logits_base = logits_ptr + q_idx * logits_q_stride + s * logits_s_stride + head_idx * logits_h_stride
        vals = tl.load(logits_base + v_offs, mask=v_mask, other=0.0)
        acc += vals * w

    # Normalize
    acc = acc / weight_sum

    # Store as bf16
    out_base = output_ptr + q_idx * num_heads * v_dim + head_idx * v_dim
    tl.store(out_base + v_offs, acc.to(tl.bfloat16), mask=v_mask)


# ========================================================================
# Cache functions
# ========================================================================

def _ensure_cache_nonpers_bf16(batch_size, kv_seq_len, total_q):
    key = ("npbf16", batch_size, kv_seq_len)
    if key in _cache:
        return _cache[key]

    nq = NUM_HEADS
    total_kv = batch_size * kv_seq_len

    kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
    num_kv_splits, num_kv_splits_indptr = get_meta_param(None, batch_size, total_kv, nq, 1, torch.bfloat16)
    o = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
    logits = torch.empty((total_q, num_kv_splits, nq, V_HEAD_DIM), dtype=torch.float32, device="cuda")
    attn_lse = torch.empty((total_q, num_kv_splits, nq, 1), dtype=torch.float32, device="cuda")

    _cache[key] = {
        "kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,
        "num_kv_splits": num_kv_splits, "num_kv_splits_indptr": num_kv_splits_indptr,
        "logits": logits, "attn_lse": attn_lse, "o": o,
    }
    return _cache[key]


def _ensure_cache_nonpers_fp8(batch_size, kv_seq_len, total_q):
    """Non-persistent fp8 with float32 logits and custom reduce."""
    key = ("npfp8_cr", batch_size, kv_seq_len)
    if key in _cache:
        return _cache[key]

    nq = NUM_HEADS
    total_kv = batch_size * kv_seq_len

    kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
    num_kv_splits, num_kv_splits_indptr = get_meta_param(None, batch_size, total_kv, nq, 1, FP8_DTYPE)

    o = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
    q_fp8 = torch.empty((total_q, nq * QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda")
    q_scale = torch.ones(1, dtype=torch.float32, device="cuda")

    # Always float32 logits — never alias to output
    logits = torch.empty((total_q, num_kv_splits, nq, V_HEAD_DIM), dtype=torch.float32, device="cuda")
    attn_lse = torch.empty((total_q, num_kv_splits, nq, 1), dtype=torch.float32, device="cuda")

    _cache[key] = {
        "kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,
        "num_kv_splits": num_kv_splits, "num_kv_splits_indptr": num_kv_splits_indptr,
        "logits": logits, "attn_lse": attn_lse, "o": o,
        "q_fp8": q_fp8, "q_scale": q_scale,
    }
    return _cache[key]


def _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, persistent_splits, fast_mode, kv_gran=16):
    key = ("pfp8", batch_size, kv_seq_len, persistent_splits, fast_mode, kv_gran)
    if key in _cache:
        return _cache[key]

    max_q_len = 1
    nq, nkv = NUM_HEADS, NUM_KV_HEADS
    total_kv = batch_size * kv_seq_len

    kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")

    info = get_mla_metadata_info_v1(
        batch_size, max_q_len, nq, FP8_DTYPE, FP8_DTYPE,
        is_sparse=False, fast_mode=fast_mode,
        num_kv_splits=persistent_splits, intra_batch_mode=True,
    )
    work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
    (work_metadata, work_indptr, work_info_set,
     reduce_indptr, reduce_final_map, reduce_partial_map) = work

    get_mla_metadata_v1(
        qo_indptr, kv_indptr, kv_last_page_len,
        nq // nkv, nkv, True,
        work_metadata, work_info_set, work_indptr,
        reduce_indptr, reduce_final_map, reduce_partial_map,
        page_size=PAGE_SIZE,
        kv_granularity=max(PAGE_SIZE, kv_gran),
        max_seqlen_qo=max_q_len,
        uni_seqlen_qo=max_q_len,
        fast_mode=fast_mode,
        max_split_per_batch=persistent_splits,
        intra_batch_mode=True,
        dtype_q=FP8_DTYPE,
        dtype_kv=FP8_DTYPE,
    )

    num_partials = reduce_partial_map.size(0)
    logits = torch.empty((num_partials, 1, nq, V_HEAD_DIM), dtype=torch.float32, device="cuda")
    attn_lse = torch.empty((num_partials, 1, nq, 1), dtype=torch.float32, device="cuda")
    o = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
    q_fp8 = torch.empty((total_q, nq * QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda")
    q_scale = torch.ones(1, dtype=torch.float32, device="cuda")

    _cache[key] = {
        "kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,
        "work_metadata": work_metadata, "work_indptr": work_indptr,
        "work_info_set": work_info_set, "reduce_indptr": reduce_indptr,
        "reduce_final_map": reduce_final_map, "reduce_partial_map": reduce_partial_map,
        "logits": logits, "attn_lse": attn_lse, "o": o,
        "q_fp8": q_fp8, "q_scale": q_scale,
        "num_partials": num_partials,
    }
    return _cache[key]


# ========================================================================
# Main dispatch
# ========================================================================

def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data

    batch_size = config["batch_size"]
    kv_seq_len = config["kv_seq_len"]
    total_q = q.shape[0]
    total_kv = batch_size * kv_seq_len

    # ---- Tier 1: batch<=4 -> bf16/bf16 non-persistent (s1, s2) ----
    if batch_size <= 4:
        kv_bf16 = kv_data["bf16"]
        kv_4d = kv_bf16.view(total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)
        c = _ensure_cache_nonpers_bf16(batch_size, kv_seq_len, total_q)

        aiter.mla_decode_stage1_asm_fwd(
            q.view(-1, NUM_HEADS, QK_HEAD_DIM), kv_4d,
            qo_indptr, kv_indptr, c["kv_indices"], c["kv_last_page_len"],
            c["num_kv_splits_indptr"],
            None, None, None,
            1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
            c["logits"], c["attn_lse"], c["o"],
            None, None,
        )

        # Use custom reduce for bf16 nonpers too (for consistency)
        num_splits = c["num_kv_splits"]
        BLOCK_V = 128
        n_v_blocks = triton.cdiv(V_HEAD_DIM, BLOCK_V)
        _custom_reduce_kernel[(total_q, NUM_HEADS, n_v_blocks)](
            c["logits"], c["attn_lse"], c["o"],
            total_q,
            num_splits=num_splits,
            num_heads=NUM_HEADS,
            v_dim=V_HEAD_DIM,
            BLOCK_V=BLOCK_V,
            num_warps=4,
        )
        return c["o"]

    # ---- Tier 2: nonpers fp8 with custom reduce (s3, s5 ONLY) ----
    # Only kv<=1024 is safe for nonpers in leaderboard mode.
    # kv=8192 nonpers FAILS leaderboard (v94 failed on s4).
    elif kv_seq_len <= 1024 and batch_size <= 64:
        kv_buffer_fp8, kv_scale = kv_data["fp8"]
        kv_buffer_4d = kv_buffer_fp8.view(total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)

        c = _ensure_cache_nonpers_fp8(batch_size, kv_seq_len, total_q)

        # Fast FP8 quant: copy_ cast (scale=1.0)
        q_2d = q.view(total_q, NUM_HEADS * QK_HEAD_DIM)
        c["q_fp8"].copy_(q_2d)

        # Stage 1: non-persistent fp8 attention
        aiter.mla_decode_stage1_asm_fwd(
            c["q_fp8"].view(-1, NUM_HEADS, QK_HEAD_DIM), kv_buffer_4d,
            qo_indptr, kv_indptr, c["kv_indices"], c["kv_last_page_len"],
            c["num_kv_splits_indptr"],
            None, None, None,
            1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
            c["logits"], c["attn_lse"], c["o"],
            c["q_scale"], kv_scale,
        )

        # Stage 2: CUSTOM reduce (replaces aiter's _fwd_kernel_stage2_asm)
        num_splits = c["num_kv_splits"]
        BLOCK_V = 128
        n_v_blocks = triton.cdiv(V_HEAD_DIM, BLOCK_V)
        _custom_reduce_kernel[(total_q, NUM_HEADS, n_v_blocks)](
            c["logits"], c["attn_lse"], c["o"],
            total_q,
            num_splits=num_splits,
            num_heads=NUM_HEADS,
            v_dim=V_HEAD_DIM,
            BLOCK_V=BLOCK_V,
            num_warps=4,
        )
        return c["o"]

    # ---- Tier 3: fp8 persistent (s4, s6, s7, s8) ----
    else:
        kv_buffer_fp8, kv_scale = kv_data["fp8"]
        kv_buffer_4d = kv_buffer_fp8.view(total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)

        # Per-shape split tuning
        if total_kv >= 1000000:
            splits, fast_mode = 4, False      # s8
        elif total_kv >= 300000:
            splits, fast_mode = 8, False      # s6
        elif batch_size >= 256:
            splits, fast_mode = 1, False      # s7
        elif batch_size >= 64:
            splits, fast_mode = 4, False      # s5 won't reach here (kv<=1024 → tier 2)
        else:
            splits, fast_mode = 32, True      # s4

        kv_gran = 64  # use 64 for all persistent shapes (was 16 for kv<=1024)
        c = _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, splits, fast_mode, kv_gran)

        q_2d = q.view(total_q, NUM_HEADS * QK_HEAD_DIM)
        c["q_fp8"].copy_(q_2d)

        aiter.mla_decode_stage1_asm_fwd(
            c["q_fp8"].view(-1, NUM_HEADS, QK_HEAD_DIM), kv_buffer_4d,
            qo_indptr, kv_indptr, c["kv_indices"], c["kv_last_page_len"],
            None, c["work_metadata"], c["work_indptr"], c["work_info_set"],
            1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
            c["logits"], c["attn_lse"], c["o"],
            c["q_scale"], kv_scale,
        )

        aiter.mla_reduce_v1(
            c["logits"], c["attn_lse"],
            c["reduce_indptr"], c["reduce_final_map"], c["reduce_partial_map"],
            1, c["o"], None,
        )
        return c["o"]
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