Skip to content
KernelIndex
Search⌘K

submission 753817

sean_nobricks · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:cc23736b2700f4b273bb18e9aafd822bb60526b8302904ef4aafd430b725da97
license declaredunknown
license concludedunknown
authorssean_nobricks
imported2026-08-15

Techniques

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

mmaqk = tl.dot(q_nope, tl.trans(k_nope))
num-warps = 4num_warps=4, num_stages=2, waves_per_eu=1,
online-softmaxm_i_new = tl.maximum(m_i, tl.max(qk, 1))
persistent-kernelper-tile mask. An AITER non-persistent fallback handles very large workloads
split-k"""MLA decode — Triton fp8 flash-decoding with split-K and LSE reduction.
stages = 2num_warps=4, num_stages=2, waves_per_eu=1,
tile-n = 16BLOCK_N = 16

Kernel source

submission.py473 lines
"""MLA decode — Triton fp8 flash-decoding with split-K and LSE reduction.

Two-stage flash-decoding for DeepSeek R1 forward_absorb MLA on MI355X:
- Stage 1 splits the KV sequence across CTAs, runs Q@K^T + online softmax + P@V
  with fp8 KV loads and 16-head MQA packing, and writes per-split partial
  outputs and LSEs.
- Stage 2 reduces the per-split partials with an LSE-weighted sum.

An exact-no-mask stage 1 specialization is dispatched when the KV length
divides evenly into the chosen split count and BLOCK_N, eliminating the
per-tile mask. An AITER non-persistent fallback handles very large workloads
where library ASM bandwidth exceeds the custom path.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t

# DeepSeek R1 forward_absorb MLA constants
NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576       # kv_lora_rank (512) + qk_rope_head_dim (64)
V_HEAD_DIM = 512        # = kv_lora_rank
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
LOG2E = 1.44269504

from aiter import dtypes as aiter_dtypes
FP8_DTYPE = aiter_dtypes.fp8

# Pre-allocated workspace buffers
_workspace_partial_out = None
_workspace_partial_lse = None
_workspace_final_out = None


# ---------------------------------------------------------------------------
# Stage 1: flash-decoding with fp8 KV loads and MQA head packing
# ---------------------------------------------------------------------------

@triton.jit
def mla_flash_decode_stage1(
    Q, KV_fp8, Out_partial, LSE_partial,
    qo_indptr, kv_indptr,
    KV_scale_ptr,
    sm_scale_log2e,
    stride_qt, stride_qh, stride_qd,
    stride_kvt, stride_kvd,
    stride_opt_b, stride_opt_s, stride_opt_h, stride_opt_d,
    stride_lse_b, stride_lse_s, stride_lse_h,
    BLOCK_H: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_DK: tl.constexpr,
    BLOCK_DPE: tl.constexpr,
    BLOCK_DV: tl.constexpr,
    NUM_KV_SPLITS: tl.constexpr,
):
    batch_id = tl.program_id(0)
    split_id = tl.program_id(1)

    tl.assume(stride_qt > 0)
    tl.assume(stride_qh > 0)
    tl.assume(stride_qd > 0)
    tl.assume(stride_kvt > 0)
    tl.assume(stride_kvd > 0)
    tl.assume(stride_opt_b > 0)
    tl.assume(stride_opt_s > 0)
    tl.assume(stride_opt_h > 0)
    tl.assume(stride_opt_d > 0)
    tl.assume(stride_lse_b > 0)
    tl.assume(stride_lse_s > 0)
    tl.assume(stride_lse_h > 0)

    q_start = tl.load(qo_indptr + batch_id)
    kv_start = tl.load(kv_indptr + batch_id)
    kv_end = tl.load(kv_indptr + batch_id + 1)
    kv_len = kv_end - kv_start

    kv_per_split = tl.cdiv(kv_len, NUM_KV_SPLITS)
    split_kv_start = split_id * kv_per_split
    split_kv_end = tl.minimum(split_kv_start + kv_per_split, kv_len)

    # Early exit for unused splits
    if split_kv_start >= kv_len:
        offs_h = tl.arange(0, BLOCK_H)
        offs_dv = tl.arange(0, BLOCK_DV)
        out_ptrs = (Out_partial + batch_id * stride_opt_b + split_id * stride_opt_s
                    + offs_h[:, None] * stride_opt_h + offs_dv[None, :] * stride_opt_d)
        tl.store(out_ptrs, tl.zeros([BLOCK_H, BLOCK_DV], dtype=tl.bfloat16))
        lse_ptrs = (LSE_partial + batch_id * stride_lse_b + split_id * stride_lse_s
                    + offs_h * stride_lse_h)
        tl.store(lse_ptrs, tl.full([BLOCK_H], value=float('-inf'), dtype=tl.float32))
        return

    kv_scale = tl.load(KV_scale_ptr).to(tl.float32)
    qk_scale = kv_scale * sm_scale_log2e

    offs_h = tl.arange(0, BLOCK_H)
    offs_dk = tl.max_contiguous(tl.multiple_of(tl.arange(0, BLOCK_DK), BLOCK_DK), BLOCK_DK)
    offs_dpe = tl.max_contiguous(tl.multiple_of(tl.arange(0, BLOCK_DPE), BLOCK_DPE), BLOCK_DPE)

    q_base = Q + q_start * stride_qt
    q_nope = tl.load(
        q_base + offs_h[:, None] * stride_qh + offs_dk[None, :] * stride_qd,
        cache_modifier=".cg",
    )
    q_pe = tl.load(
        q_base + offs_h[:, None] * stride_qh + (BLOCK_DK + offs_dpe[None, :]) * stride_qd,
        cache_modifier=".cg",
    )
    q_nope = (q_nope.to(tl.float32) * qk_scale).to(tl.bfloat16)
    q_pe = (q_pe.to(tl.float32) * qk_scale).to(tl.bfloat16)

    m_i = tl.full([BLOCK_H], value=float('-inf'), dtype=tl.float32)
    l_i = tl.zeros([BLOCK_H], dtype=tl.float32)
    acc = tl.zeros([BLOCK_H, BLOCK_DV], dtype=tl.float32)

    offs_n = tl.arange(0, BLOCK_N)
    kv_base = KV_fp8 + kv_start * stride_kvt

    for n_start in range(split_kv_start, split_kv_end, BLOCK_N):
        n_offs = n_start + offs_n
        kv_mask = n_offs < split_kv_end

        kv_ptrs_base = kv_base + n_offs[:, None] * stride_kvt
        k_nope_fp8 = tl.load(
            kv_ptrs_base + offs_dk[None, :] * stride_kvd,
            mask=kv_mask[:, None], other=0.0,
            cache_modifier=".cg",
        )
        k_pe_fp8 = tl.load(
            kv_ptrs_base + (BLOCK_DK + offs_dpe[None, :]) * stride_kvd,
            mask=kv_mask[:, None], other=0.0,
            cache_modifier=".cg",
        )

        k_nope = k_nope_fp8.to(tl.bfloat16)
        k_pe = k_pe_fp8.to(tl.bfloat16)

        qk = tl.dot(q_nope, tl.trans(k_nope))
        qk += tl.dot(q_pe, tl.trans(k_pe))

        qk = tl.where(kv_mask[None, :], qk, float('-inf'))

        m_i_new = tl.maximum(m_i, tl.max(qk, 1))
        alpha = tl.math.exp2(m_i - m_i_new)
        p = tl.math.exp2(qk - m_i_new[:, None])

        l_i = l_i * alpha + tl.sum(p, 1)
        acc = acc * alpha[:, None]

        acc += tl.dot(p.to(tl.bfloat16), k_nope)
        m_i = m_i_new

    acc = acc * kv_scale / l_i[:, None]

    offs_dv = tl.arange(0, BLOCK_DV)
    out_ptrs = (Out_partial + batch_id * stride_opt_b + split_id * stride_opt_s
                + offs_h[:, None] * stride_opt_h + offs_dv[None, :] * stride_opt_d)
    tl.store(out_ptrs, acc.to(tl.bfloat16))

    lse = (tl.math.log2(l_i) + m_i) * 0.6931471805599453
    lse_ptrs = (LSE_partial + batch_id * stride_lse_b + split_id * stride_lse_s
                + offs_h * stride_lse_h)
    tl.store(lse_ptrs, lse)


@triton.jit
def mla_flash_decode_stage1_exact_nomask(
    Q, KV_fp8, Out_partial, LSE_partial,
    qo_indptr, kv_indptr,
    KV_scale_ptr,
    sm_scale_log2e,
    stride_qt, stride_qh, stride_qd,
    stride_kvt, stride_kvd,
    stride_opt_b, stride_opt_s, stride_opt_h, stride_opt_d,
    stride_lse_b, stride_lse_s, stride_lse_h,
    BLOCK_H: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_DK: tl.constexpr,
    BLOCK_DPE: tl.constexpr,
    BLOCK_DV: tl.constexpr,
    NUM_KV_SPLITS: tl.constexpr,
):
    batch_id = tl.program_id(0)
    split_id = tl.program_id(1)

    tl.assume(stride_qt > 0)
    tl.assume(stride_qh > 0)
    tl.assume(stride_qd > 0)
    tl.assume(stride_kvt > 0)
    tl.assume(stride_kvd > 0)
    tl.assume(stride_opt_b > 0)
    tl.assume(stride_opt_s > 0)
    tl.assume(stride_opt_h > 0)
    tl.assume(stride_opt_d > 0)
    tl.assume(stride_lse_b > 0)
    tl.assume(stride_lse_s > 0)
    tl.assume(stride_lse_h > 0)

    q_start = tl.load(qo_indptr + batch_id)
    kv_start = tl.load(kv_indptr + batch_id)
    kv_end = tl.load(kv_indptr + batch_id + 1)
    kv_len = kv_end - kv_start

    split_kv_tokens = kv_len // NUM_KV_SPLITS
    split_kv_tokens = tl.multiple_of(split_kv_tokens, BLOCK_N)
    tl.assume(split_kv_tokens >= BLOCK_N)
    split_kv_start = split_id * split_kv_tokens
    split_kv_start = tl.multiple_of(split_kv_start, BLOCK_N)
    split_kv_end = split_kv_start + split_kv_tokens

    kv_scale = tl.load(KV_scale_ptr).to(tl.float32)
    qk_scale = kv_scale * sm_scale_log2e

    offs_h = tl.arange(0, BLOCK_H)
    offs_dk = tl.max_contiguous(tl.multiple_of(tl.arange(0, BLOCK_DK), BLOCK_DK), BLOCK_DK)
    offs_dpe = tl.max_contiguous(tl.multiple_of(tl.arange(0, BLOCK_DPE), BLOCK_DPE), BLOCK_DPE)

    q_base = Q + q_start * stride_qt
    q_nope = tl.load(
        q_base + offs_h[:, None] * stride_qh + offs_dk[None, :] * stride_qd,
        cache_modifier=".cg",
    )
    q_pe = tl.load(
        q_base + offs_h[:, None] * stride_qh + (BLOCK_DK + offs_dpe[None, :]) * stride_qd,
        cache_modifier=".cg",
    )
    q_nope = (q_nope.to(tl.float32) * qk_scale).to(tl.bfloat16)
    q_pe = (q_pe.to(tl.float32) * qk_scale).to(tl.bfloat16)

    m_i = tl.full([BLOCK_H], value=float("-inf"), dtype=tl.float32)
    l_i = tl.zeros([BLOCK_H], dtype=tl.float32)
    acc = tl.zeros([BLOCK_H, BLOCK_DV], dtype=tl.float32)

    offs_n = tl.arange(0, BLOCK_N)
    kv_base = KV_fp8 + kv_start * stride_kvt

    for n_start in range(split_kv_start, split_kv_end, BLOCK_N):
        n_offs = n_start + offs_n
        kv_ptrs_base = kv_base + n_offs[:, None] * stride_kvt
        k_nope_fp8 = tl.load(
            kv_ptrs_base + offs_dk[None, :] * stride_kvd,
            cache_modifier=".cg",
        )
        k_pe_fp8 = tl.load(
            kv_ptrs_base + (BLOCK_DK + offs_dpe[None, :]) * stride_kvd,
            cache_modifier=".cg",
        )

        k_nope = k_nope_fp8.to(tl.bfloat16)
        k_pe = k_pe_fp8.to(tl.bfloat16)

        qk = tl.dot(q_nope, tl.trans(k_nope))
        qk += tl.dot(q_pe, tl.trans(k_pe))

        m_i_new = tl.maximum(m_i, tl.max(qk, 1))
        alpha = tl.math.exp2(m_i - m_i_new)
        p = tl.math.exp2(qk - m_i_new[:, None])

        l_i = l_i * alpha + tl.sum(p, 1)
        acc = acc * alpha[:, None]
        acc += tl.dot(p.to(tl.bfloat16), k_nope)
        m_i = m_i_new

    acc = acc * kv_scale / l_i[:, None]

    offs_dv = tl.arange(0, BLOCK_DV)
    out_ptrs = (Out_partial + batch_id * stride_opt_b + split_id * stride_opt_s
                + offs_h[:, None] * stride_opt_h + offs_dv[None, :] * stride_opt_d)
    tl.store(out_ptrs, acc.to(tl.bfloat16))

    lse = (tl.math.log2(l_i) + m_i) * 0.6931471805599453
    lse_ptrs = (LSE_partial + batch_id * stride_lse_b + split_id * stride_lse_s
                + offs_h * stride_lse_h)
    tl.store(lse_ptrs, lse)


# ---------------------------------------------------------------------------
# Stage 2: LSE-weighted reduction of split partial outputs
# ---------------------------------------------------------------------------

@triton.jit
def mla_flash_decode_stage2(
    Out_partial, LSE_partial, Out_final,
    stride_opt_b, stride_opt_s, stride_opt_h, stride_opt_d,
    stride_lse_b, stride_lse_s, stride_lse_h,
    stride_of_t, stride_of_h, stride_of_d,
    BLOCK_DV: tl.constexpr,
    NUM_KV_SPLITS: tl.constexpr,
):
    batch_id = tl.program_id(0)
    head_id = tl.program_id(1)

    lse_base = LSE_partial + batch_id * stride_lse_b + head_id * stride_lse_h
    out_base = Out_partial + batch_id * stride_opt_b + head_id * stride_opt_h

    lse_max = tl.load(lse_base)
    for s in tl.static_range(1, NUM_KV_SPLITS):
        lse_s = tl.load(lse_base + s * stride_lse_s)
        lse_max = tl.where(lse_s > lse_max, lse_s, lse_max)

    offs_dv = tl.arange(0, BLOCK_DV)
    acc = tl.zeros([BLOCK_DV], dtype=tl.float32)
    weight_sum = 0.0

    for s in tl.static_range(NUM_KV_SPLITS):
        lse_s = tl.load(lse_base + s * stride_lse_s)
        w = tl.exp(lse_s - lse_max)
        partial = tl.load(out_base + s * stride_opt_s + offs_dv * stride_opt_d).to(tl.float32)
        acc += w * partial
        weight_sum += w

    acc = acc / weight_sum

    out_ptrs = (Out_final + batch_id * stride_of_t + head_id * stride_of_h
                + offs_dv * stride_of_d)
    tl.store(out_ptrs, acc.to(tl.bfloat16))


def _quantize_fp8(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)


_cached_aiter_data = {}


def _library_mla_decode_nonpersistent(q, kv_data, qo_indptr, kv_indptr, config):
    global _cached_aiter_data
    from aiter.mla import mla_decode_fwd
    q_fp8, q_scale = _quantize_fp8(q)
    kv_buffer_fp8, kv_scale = kv_data["fp8"]
    batch_size = config["batch_size"]
    kv_seq_len = config["kv_seq_len"]
    total_kv_len = batch_size * kv_seq_len
    kv_buffer_4d = kv_buffer_fp8.view(
        kv_buffer_fp8.shape[0], PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM
    )
    cache_key = (batch_size, kv_seq_len)
    if cache_key not in _cached_aiter_data:
        kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
        kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
        o = torch.empty((q.shape[0], NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
        _cached_aiter_data[cache_key] = (kv_indices, kv_last_page_len, o)
    kv_indices, kv_last_page_len, o = _cached_aiter_data[cache_key]
    mla_decode_fwd(
        q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM),
        kv_buffer_4d, o,
        qo_indptr, kv_indptr, kv_indices, kv_last_page_len,
        config["q_seq_len"],
        page_size=PAGE_SIZE, nhead_kv=NUM_KV_HEADS,
        sm_scale=SM_SCALE, logit_cap=0.0,
        q_scale=q_scale, kv_scale=kv_scale,
    )
    return o


def _choose_kv_splits(batch_size, kv_len, block_n=64, num_cus=256):
    """Choose splits to balance CU utilization vs stage2 overhead."""
    total_ctas = batch_size
    if total_ctas >= num_cus:
        return 1
    target_cu = max(1, num_cus // batch_size)
    max_useful = max(1, kv_len // (4 * block_n))
    splits = min(target_cu, max_useful)
    for po2 in [1, 2, 4, 8, 16, 32, 64]:
        if po2 >= splits:
            return po2
    return 64


# ---------------------------------------------------------------------------
# Entry point
# ---------------------------------------------------------------------------

def custom_kernel(data: input_t) -> output_t:
    global _workspace_partial_out, _workspace_partial_lse, _workspace_final_out

    q, kv_data, qo_indptr, kv_indptr, config = data

    batch_size = config["batch_size"]
    kv_seq_len = config["kv_seq_len"]

    # AITER fallback for very large workloads where library ASM bandwidth
    # exceeds the custom path; bs=256/kv=8192 stays on the custom kernel.
    total_kv_tokens = batch_size * kv_seq_len
    use_aiter_fallback = total_kv_tokens >= 1048576 and not (batch_size == 256 and kv_seq_len == 8192)
    if use_aiter_fallback:
        return _library_mla_decode_nonpersistent(q, kv_data, qo_indptr, kv_indptr, config)

    kv_buffer_fp8, kv_scale_tensor = kv_data["fp8"]

    BLOCK_H = 16
    if batch_size == 4 and kv_seq_len == 1024:
        BLOCK_N = 16
    elif batch_size <= 32 and kv_seq_len <= 1024:
        BLOCK_N = 32
    else:
        BLOCK_N = 64
    BLOCK_DK = 512
    BLOCK_DPE = 64
    BLOCK_DV = 512
    NUM_KV_SPLITS = _choose_kv_splits(batch_size, kv_seq_len, BLOCK_N)
    uses_exact_split_stage1 = (kv_seq_len % (NUM_KV_SPLITS * BLOCK_N) == 0)

    po_shape = (batch_size, NUM_KV_SPLITS, NUM_HEADS, V_HEAD_DIM)
    pl_shape = (batch_size, NUM_KV_SPLITS, NUM_HEADS)
    if _workspace_partial_out is None or _workspace_partial_out.shape != po_shape:
        _workspace_partial_out = torch.empty(po_shape, dtype=torch.bfloat16, device=q.device)
        _workspace_partial_lse = torch.empty(pl_shape, dtype=torch.float32, device=q.device)

    sm_scale_log2e = SM_SCALE * LOG2E
    stride_kvt = kv_buffer_fp8.stride(0)
    stride_kvd = kv_buffer_fp8.stride(2)

    grid_stage1 = (batch_size, NUM_KV_SPLITS)
    stage1_args = (
        q, kv_buffer_fp8, _workspace_partial_out, _workspace_partial_lse,
        qo_indptr, kv_indptr,
        kv_scale_tensor,
        sm_scale_log2e,
        q.stride(0), q.stride(1), q.stride(2),
        stride_kvt, stride_kvd,
        _workspace_partial_out.stride(0), _workspace_partial_out.stride(1),
        _workspace_partial_out.stride(2), _workspace_partial_out.stride(3),
        _workspace_partial_lse.stride(0), _workspace_partial_lse.stride(1),
        _workspace_partial_lse.stride(2),
    )
    stage1_meta = dict(
        BLOCK_H=BLOCK_H, BLOCK_N=BLOCK_N,
        BLOCK_DK=BLOCK_DK, BLOCK_DPE=BLOCK_DPE, BLOCK_DV=BLOCK_DV,
        NUM_KV_SPLITS=NUM_KV_SPLITS,
        num_warps=4, num_stages=2, waves_per_eu=1,
        schedule_hint="memory-bound-attention",
    )
    if uses_exact_split_stage1:
        mla_flash_decode_stage1_exact_nomask[grid_stage1](
            *stage1_args,
            **stage1_meta,
        )
    else:
        mla_flash_decode_stage1[grid_stage1](
            *stage1_args,
            **stage1_meta,
        )

    if NUM_KV_SPLITS == 1:
        return _workspace_partial_out.squeeze(1)

    fo_shape = (q.shape[0], NUM_HEADS, V_HEAD_DIM)
    if _workspace_final_out is None or _workspace_final_out.shape != fo_shape:
        _workspace_final_out = torch.empty(fo_shape, dtype=torch.bfloat16, device=q.device)

    grid_stage2 = (batch_size, NUM_HEADS)
    mla_flash_decode_stage2[grid_stage2](
        _workspace_partial_out, _workspace_partial_lse, _workspace_final_out,
        _workspace_partial_out.stride(0), _workspace_partial_out.stride(1),
        _workspace_partial_out.stride(2), _workspace_partial_out.stride(3),
        _workspace_partial_lse.stride(0), _workspace_partial_lse.stride(1),
        _workspace_partial_lse.stride(2),
        _workspace_final_out.stride(0), _workspace_final_out.stride(1),
        _workspace_final_out.stride(2),
        BLOCK_DV=V_HEAD_DIM,
        NUM_KV_SPLITS=NUM_KV_SPLITS,
        num_warps=4, num_stages=1,
    )

    return _workspace_final_out
scrolls · 473 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