Skip to content
KernelIndex
Search⌘K

submission 737458

Jayluci4 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_mla_hybrid.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-737458?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
53.0µs
#174 of 766
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:264b18534353576efa1abe856907da461cb1d8bd252a1e69aee72e702895d79a
license declaredunknown
license concludedunknown
authorsJayluci4
imported2026-08-15

Techniques

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

mmascores = tl.dot(q_nope, tl.trans(kv_nope_bf))
num-warps = 4NUM_SPLITS=splits, BLOCK_KV=64, num_warps=4, num_stages=2)
online-softmaxm_new = tl.maximum(m_i, m_block)
stages = 2NUM_SPLITS=splits, BLOCK_KV=64, num_warps=4, num_stages=2)

Kernel source

submission_mla_hybrid.py191 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
Hybrid MLA decode: Triton fused for bs<=32, ASM for bs>=64.
Triton: single launch, bf16 Q (no quant), MQA, fp8 KV.
ASM: hand-tuned BW, splits=1 for bs=256 (no reduce).
"""
import gc
import sys
import os

gc.disable()
sys.setswitchinterval(1000.0)
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"

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

NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / (576 ** 0.5)
PAGE_SIZE = 1
FP8 = aiter_dtypes.fp8

# ASM only for shapes where it beats Triton (large bs × large kv)
_ASM_SPLITS = {(64, 8192): 4, (256, 8192): 1}
# Triton for everything else (kv_scale absorbed, single-launch wins)
_TRI_SPLITS = {
    (4, 1024): 64, (4, 8192): 64,
    (32, 1024): 8, (32, 8192): 8,
    (64, 1024): 4, (256, 1024): 1,
}
# Route: use ASM when shape is in _ASM_SPLITS
_USE_ASM = {(64, 8192), (256, 8192)}

_cache = {}


@triton.jit
def _mla_fused_kernel(
    Q_ptr, KV_fp8_ptr, O_ptr, Mid_O_ptr, Mid_LSE_ptr,
    kv_indptr_ptr, kv_scale_ptr, sm_scale,
    NUM_SPLITS: tl.constexpr, BLOCK_KV: tl.constexpr,
):
    pid = tl.program_id(0)
    batch_idx = pid // NUM_SPLITS
    split_idx = pid % NUM_SPLITS
    kv_start = tl.load(kv_indptr_ptr + batch_idx)
    kv_end = tl.load(kv_indptr_ptr + batch_idx + 1)
    kv_len = kv_end - kv_start
    split_size = (kv_len + NUM_SPLITS - 1) // NUM_SPLITS
    s_start = kv_start + split_idx * split_size
    s_end = tl.minimum(s_start + split_size, kv_end)
    kv_scale = tl.load(kv_scale_ptr)
    sm_scale_adj = sm_scale * kv_scale
    q_base = batch_idx * 16 * 576
    offs_h = tl.arange(0, 16)
    offs_512 = tl.arange(0, 512)
    offs_64 = tl.arange(0, 64)
    q_nope = tl.load(Q_ptr + q_base + offs_h[:, None] * 576 + offs_512[None, :]).to(tl.bfloat16)
    q_rope = tl.load(Q_ptr + q_base + offs_h[:, None] * 576 + 512 + offs_64[None, :]).to(tl.bfloat16)
    m_i = tl.full([16], float("-inf"), dtype=tl.float32)
    l_i = tl.zeros([16], dtype=tl.float32)
    o_acc = tl.zeros([16, 512], dtype=tl.float32)
    offs_kv = tl.arange(0, BLOCK_KV)
    for kv_off in range(0, split_size, BLOCK_KV):
        kv_pos = s_start + kv_off
        valid = (kv_pos + offs_kv) < s_end
        kv_rows = tl.minimum(kv_pos + offs_kv, s_end - 1)
        kv_nope_bf = tl.load(KV_fp8_ptr + kv_rows[:, None] * 576 + offs_512[None, :]).to(tl.bfloat16)
        kv_rope_bf = tl.load(KV_fp8_ptr + kv_rows[:, None] * 576 + 512 + offs_64[None, :]).to(tl.bfloat16)
        scores = tl.dot(q_nope, tl.trans(kv_nope_bf))
        scores += tl.dot(q_rope, tl.trans(kv_rope_bf))
        scores = scores.to(tl.float32) * sm_scale_adj
        scores = tl.where(valid[None, :], scores, float("-inf"))
        m_block = tl.max(scores, axis=1)
        m_new = tl.maximum(m_i, m_block)
        alpha = tl.exp(m_i - m_new)
        p = tl.exp(scores - m_new[:, None])
        l_new = l_i * alpha + tl.sum(p, axis=1)
        p_bf = p.to(tl.bfloat16)
        o_acc = o_acc * alpha[:, None] + tl.dot(p_bf, kv_nope_bf).to(tl.float32)
        m_i = m_new
        l_i = l_new
    safe_l = tl.where(l_i > 0, l_i, 1.0)
    o_acc = (o_acc / safe_l[:, None]) * kv_scale
    if NUM_SPLITS == 1:
        o_base = batch_idx * 16 * 512
        tl.store(O_ptr + o_base + offs_h[:, None] * 512 + offs_512[None, :], o_acc.to(tl.bfloat16))
    else:
        mid_idx = batch_idx * NUM_SPLITS + split_idx
        mid_base = mid_idx * 16 * 512
        tl.store(Mid_O_ptr + mid_base + offs_h[:, None] * 512 + offs_512[None, :], o_acc)
        lse = tl.where(l_i > 0, m_i + tl.log(l_i), float("-inf"))
        tl.store(Mid_LSE_ptr + mid_idx * 16 + offs_h, lse)


@triton.jit
def _reduce_kernel(Mid_O_ptr, Mid_LSE_ptr, O_ptr, NUM_SPLITS: tl.constexpr):
    pid_b = tl.program_id(0)
    pid_h = tl.program_id(1)
    offs_v = tl.arange(0, 512)
    m_max = tl.full([1], float("-inf"), dtype=tl.float32)
    for s in range(NUM_SPLITS):
        lse = tl.load(Mid_LSE_ptr + (pid_b * NUM_SPLITS + s) * 16 + pid_h)
        m_max = tl.maximum(m_max, lse)
    acc = tl.zeros([512], dtype=tl.float32)
    l_sum = tl.zeros([1], dtype=tl.float32)
    for s in range(NUM_SPLITS):
        idx = pid_b * NUM_SPLITS + s
        lse = tl.load(Mid_LSE_ptr + idx * 16 + pid_h)
        w = tl.exp(lse - m_max)
        l_sum += w
        partial = tl.load(Mid_O_ptr + idx * 16 * 512 + pid_h * 512 + offs_v)
        acc += w * partial
    safe_l = tl.where(l_sum > 0.0, l_sum, 1.0)
    tl.store(O_ptr + pid_b * 16 * 512 + pid_h * 512 + offs_v, (acc / safe_l).to(tl.bfloat16))


def _build_triton(bs, kv_len):
    splits = _TRI_SPLITS.get((bs, kv_len), 8)
    output = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
    mid_o = torch.empty((bs * splits, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device="cuda") if splits > 1 else None
    mid_lse = torch.empty((bs * splits, NUM_HEADS), dtype=torch.float32, device="cuda") if splits > 1 else None
    return ("t", splits, output, mid_o, mid_lse)


def _build_asm(bs, kv_len):
    splits = _ASM_SPLITS.get((bs, kv_len), 1)
    total_kv = bs * kv_len
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
    kv_last_page_len = torch.ones(bs, dtype=torch.int32, device="cuda")
    q_fp8 = torch.empty((bs, NUM_HEADS, QK_HEAD_DIM), dtype=FP8, device="cuda")
    q_scale = torch.ones(1, dtype=torch.float32, device="cuda")
    output = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
    splits_indptr = torch.arange(0, (bs + 1) * splits, splits, dtype=torch.int32, device="cuda")
    if splits == 1:
        logits = output.view(bs, 1, NUM_HEADS, V_HEAD_DIM)
        attn_lse = torch.empty((bs, 1, NUM_HEADS, 1), dtype=torch.float32, device="cuda")
    else:
        logits = torch.empty((bs, splits, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device="cuda")
        attn_lse = torch.empty((bs, splits, NUM_HEADS, 1), dtype=torch.float32, device="cuda")
    return ("a", splits, kv_indices, kv_last_page_len, q_fp8, q_scale,
            output, splits_indptr, logits, attn_lse, total_kv)


def custom_kernel(data: input_t) -> output_t:
    cfg = data[4]
    bs = cfg["batch_size"]
    kv_len = cfg["kv_seq_len"]
    key = (bs, kv_len)
    if key not in _cache:
        _cache[key] = _build_asm(bs, kv_len) if key in _USE_ASM else _build_triton(bs, kv_len)
    c = _cache[key]
    if c[0] == "t":
        _, splits, output, mid_o, mid_lse = c
        kv_fp8, kv_scale = data[1]["fp8"]
        mo = mid_o if mid_o is not None else output
        ml = mid_lse if mid_lse is not None else output
        _mla_fused_kernel[(bs * splits,)](
            data[0], kv_fp8, output, mo, ml, data[3], kv_scale, SM_SCALE,
            NUM_SPLITS=splits, BLOCK_KV=64, num_warps=4, num_stages=2)
        if splits > 1:
            _reduce_kernel[(bs, NUM_HEADS)](mid_o, mid_lse, output,
                NUM_SPLITS=splits, num_warps=4, num_stages=1)
        return output
    else:
        _, splits, kv_indices, kv_last_page_len, q_fp8, q_scale, \
            output, splits_indptr, logits, attn_lse, total_kv = c
        q_fp8.copy_(data[0])
        kv_fp8, kv_scale = data[1]["fp8"]
        aiter.mla_decode_stage1_asm_fwd(
            q_fp8, kv_fp8.view(total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM),
            data[2], data[3], kv_indices, kv_last_page_len, splits_indptr,
            None, None, None, 1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
            logits, attn_lse, output, q_scale, kv_scale)
        if splits == 1:
            return output
        reduce_o = logits.view(bs * splits, NUM_HEADS, V_HEAD_DIM)
        reduce_lse = attn_lse.view(bs * splits, NUM_HEADS)
        _reduce_kernel[(bs, NUM_HEADS)](reduce_o, reduce_lse, output,
            NUM_SPLITS=splits, num_warps=4, num_stages=1)
        return output
scrolls · 191 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