Skip to content
KernelIndex
Search⌘K

submission 703561

namebot855 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:3eca8721db6c275b3ad7a7282c696484bf54bec0ab6ead84212b99e6354017f4
license declaredunknown
license concludedunknown
authorsnamebot855
imported2026-08-15

Techniques

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

mmascores += tl.dot(q_tile.to(k_tile.dtype), tl.trans(k_tile))
num-warps = 8num_warps=8, num_stages=3,
online-softmaxm_new = tl.maximum(m_i, row_max)
persistent-kernel- S4, S6, S7, S8: aiter persistent fp8 (v143)
split-kdef _reduce_splitk_parallel_s1(
stages = 3num_warps=8, num_stages=3,

Kernel source

mla_test.py652 lines
"""
mla_test: merged best kernels from mla_experiments_1235
- S1 (batch=4, kv=1024):   triton s1_v45 parallel reduce w1
- S2 (batch=4, kv=8192):   triton s2_v57 parallel reduce w2
- S3 (batch=32, kv=1024):  triton s3_v37 serial reduce w1
- S5 (batch=64, kv=1024):  triton s5_v15 no-vtile rv512
- S4, S6, S7, S8:          aiter persistent fp8 (v143)
"""
import torch
import triton
import triton.language as tl
import aiter
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, _fwd_kernel_stage2_asm

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

_cache = {}

# ============================================================
# S1 triton kernels (batch=4, kv=1024)
# splits=8, SPLIT_LEN=128, V_BLOCK=128, parallel reduce w1
# ============================================================

S1_NUM_SPLITS = 8
S1_SPLIT_LEN = 128
S1_KV_SEQ_LEN = 1024
S1_V_BLOCK = 128


@triton.jit
def _flash_decode_fp8_s1_exact_vtile(
    Q_ptr, KV_ptr, Mid_O, Mid_lse,
    stride_kv: tl.int64, kv_scale,
    sm_scale: tl.constexpr, QK_DIM: tl.constexpr, V_DIM: tl.constexpr,
    BLOCK_K: tl.constexpr, NH: tl.constexpr, V_BLOCK_C: tl.constexpr,
    NUM_SPLITS_C: tl.constexpr, SPLIT_LEN_C: tl.constexpr, KV_SEQ_LEN_C: tl.constexpr,
):
    batch_id = tl.program_id(0)
    split_id = tl.program_id(1)
    v_block_id = tl.program_id(2)

    out_idx = batch_id * NUM_SPLITS_C + split_id
    kv_pos = batch_id * KV_SEQ_LEN_C + split_id * SPLIT_LEN_C
    tl.multiple_of(kv_pos, 128)

    h_offs = tl.arange(0, NH)
    n_offs = tl.arange(0, SPLIT_LEN_C)
    v_start = v_block_id * V_BLOCK_C
    v_offs = v_start + tl.arange(0, V_BLOCK_C)
    v_mask = v_offs < V_DIM
    q_base = batch_id * NH * QK_DIM
    tl.multiple_of(q_base, 128)
    tl.multiple_of(stride_kv, 32)
    tl.assume(stride_kv > 0)

    acc = tl.zeros([NH, V_BLOCK_C], dtype=tl.float32)
    m_i = tl.full([NH], float("-inf"), dtype=tl.float32)
    l_i = tl.zeros([NH], dtype=tl.float32)
    score_scale = sm_scale * kv_scale

    scores = tl.zeros([NH, SPLIT_LEN_C], dtype=tl.float32)
    for k_start in range(0, QK_DIM, BLOCK_K):
        k_offs = tl.arange(0, BLOCK_K)
        k_mask = (k_start + k_offs) < QK_DIM
        q_tile = tl.load(Q_ptr + q_base + h_offs[:, None] * QK_DIM + k_start + k_offs[None, :], mask=k_mask[None, :], other=0.0)
        k_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + k_start + k_offs[None, :], mask=k_mask[None, :], other=0.0)
        scores += tl.dot(q_tile.to(k_tile.dtype), tl.trans(k_tile))

    scores *= score_scale
    row_max = tl.max(scores, axis=1)
    m_new = tl.maximum(m_i, row_max)
    alpha = tl.exp(m_i - m_new)
    l_i = l_i * alpha
    exp_scores = tl.exp(scores - m_new[:, None])
    l_i += tl.sum(exp_scores, axis=1)
    acc = acc * alpha[:, None]

    v_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + v_offs[None, :], mask=v_mask[None, :], other=0.0)
    acc += tl.dot(exp_scores.to(v_tile.dtype), v_tile)
    m_i = m_new

    acc = (acc * kv_scale) / l_i[:, None]
    lse_vals = m_i + tl.log(l_i)
    tl.store(Mid_O + out_idx * NH * V_DIM + h_offs[:, None] * V_DIM + v_offs[None, :], acc.to(tl.bfloat16), mask=v_mask[None, :])
    tl.store(Mid_lse + out_idx * NH + h_offs, lse_vals)


@triton.jit
def _reduce_splitk_parallel_s1(
    Mid_O, Mid_lse, O_ptr,
    NUM_SPLITS_C: tl.constexpr, V_DIM: tl.constexpr, NH: tl.constexpr, BLOCK_V: tl.constexpr,
):
    batch_id = tl.program_id(0)
    head_id = 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
    s_offs = tl.arange(0, NUM_SPLITS_C)

    all_lse = tl.load(Mid_lse + (batch_id * NUM_SPLITS_C + s_offs) * NH + head_id)
    m_global = tl.max(all_lse, axis=0)
    alphas = tl.exp(all_lse - m_global)
    l_total = tl.sum(alphas, axis=0)
    weights = alphas / l_total

    all_partials = tl.load(
        Mid_O + (batch_id * NUM_SPLITS_C + s_offs[:, None]) * NH * V_DIM + head_id * V_DIM + v_offs[None, :],
        mask=v_mask[None, :], other=0.0,
    ).to(tl.float32)

    result = tl.sum(weights[:, None] * all_partials, axis=0)
    out_base = batch_id * NH * V_DIM + head_id * V_DIM
    tl.store(O_ptr + out_base + v_offs, result.to(tl.bfloat16), mask=v_mask)


# ============================================================
# S2 triton kernels (batch=4, kv=8192)
# splits=16, SPLIT_LEN=512, V_BLOCK=128, parallel reduce w2
# ============================================================

S2_NUM_SPLITS = 16
S2_SPLIT_LEN = 512
S2_KV_SEQ_LEN = 8192
S2_V_BLOCK = 128


@triton.jit
def _flash_decode_fp8_s2_exact_vtile(
    Q_ptr, KV_ptr, Mid_O, Mid_lse,
    stride_kv: tl.int64, kv_scale,
    sm_scale: tl.constexpr, QK_DIM: tl.constexpr, V_DIM: tl.constexpr,
    BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, NH: tl.constexpr,
    V_BLOCK_C: tl.constexpr, NUM_SPLITS_C: tl.constexpr,
    SPLIT_LEN_C: tl.constexpr, KV_SEQ_LEN_C: tl.constexpr,
):
    batch_id = tl.program_id(0)
    split_id = tl.program_id(1)
    v_block_id = tl.program_id(2)

    out_idx = batch_id * NUM_SPLITS_C + split_id
    kv_start = batch_id * KV_SEQ_LEN_C + split_id * SPLIT_LEN_C
    tl.multiple_of(kv_start, 256)
    h_offs = tl.arange(0, NH)
    v_start = v_block_id * V_BLOCK_C
    v_offs = v_start + tl.arange(0, V_BLOCK_C)
    v_mask = v_offs < V_DIM
    q_base = batch_id * NH * QK_DIM
    tl.multiple_of(q_base, 128)
    tl.multiple_of(stride_kv, 32)
    tl.assume(stride_kv > 0)

    acc = tl.zeros([NH, V_BLOCK_C], dtype=tl.float32)
    m_i = tl.full([NH], float("-inf"), dtype=tl.float32)
    l_i = tl.zeros([NH], dtype=tl.float32)
    score_scale = sm_scale * kv_scale

    for kv_offset in range(0, SPLIT_LEN_C, BLOCK_N):
        kv_pos = kv_start + kv_offset
        tl.multiple_of(kv_pos, 128)
        n_offs = tl.arange(0, BLOCK_N)

        scores = tl.zeros([NH, BLOCK_N], dtype=tl.float32)
        for k_start in range(0, QK_DIM, BLOCK_K):
            k_offs = tl.arange(0, BLOCK_K)
            k_mask = (k_start + k_offs) < QK_DIM
            q_tile = tl.load(Q_ptr + q_base + h_offs[:, None] * QK_DIM + k_start + k_offs[None, :], mask=k_mask[None, :], other=0.0)
            k_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + k_start + k_offs[None, :], mask=k_mask[None, :], other=0.0)
            scores += tl.dot(q_tile.to(k_tile.dtype), tl.trans(k_tile))

        scores *= score_scale
        row_max = tl.max(scores, axis=1)
        m_new = tl.maximum(m_i, row_max)
        alpha = tl.exp(m_i - m_new)
        l_i = l_i * alpha
        exp_scores = tl.exp(scores - m_new[:, None])
        l_i += tl.sum(exp_scores, axis=1)
        acc = acc * alpha[:, None]

        v_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + v_offs[None, :], mask=v_mask[None, :], other=0.0)
        acc += tl.dot(exp_scores.to(v_tile.dtype), v_tile)
        m_i = m_new

    acc = (acc * kv_scale) / l_i[:, None]
    lse_vals = m_i + tl.log(l_i)
    tl.store(Mid_O + out_idx * NH * V_DIM + h_offs[:, None] * V_DIM + v_offs[None, :], acc.to(tl.bfloat16), mask=v_mask[None, :])
    tl.store(Mid_lse + out_idx * NH + h_offs, lse_vals)


@triton.jit
def _reduce_splitk_parallel_s2(
    Mid_O, Mid_lse, O_ptr,
    NUM_SPLITS_C: tl.constexpr, V_DIM: tl.constexpr, NH: tl.constexpr, BLOCK_V: tl.constexpr,
):
    batch_id = tl.program_id(0)
    head_id = 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
    s_offs = tl.arange(0, NUM_SPLITS_C)

    all_lse = tl.load(Mid_lse + (batch_id * NUM_SPLITS_C + s_offs) * NH + head_id)
    m_global = tl.max(all_lse, axis=0)
    alphas = tl.exp(all_lse - m_global)
    l_total = tl.sum(alphas, axis=0)
    weights = alphas / l_total

    all_partials = tl.load(
        Mid_O + (batch_id * NUM_SPLITS_C + s_offs[:, None]) * NH * V_DIM + head_id * V_DIM + v_offs[None, :],
        mask=v_mask[None, :], other=0.0,
    ).to(tl.float32)

    result = tl.sum(weights[:, None] * all_partials, axis=0)
    out_base = batch_id * NH * V_DIM + head_id * V_DIM
    tl.store(O_ptr + out_base + v_offs, result.to(tl.bfloat16), mask=v_mask)


# ============================================================
# S3 triton kernels (batch=32, kv=1024)
# splits=4, SPLIT_LEN=256, V_BLOCK=256, serial reduce w1
# ============================================================

S3_NUM_SPLITS = 4
S3_SPLIT_LEN = 256
S3_KV_SEQ_LEN = 1024
S3_V_BLOCK = 256


@triton.jit
def _flash_decode_fp8_s3_exact_vtile(
    Q_ptr, KV_ptr, Mid_O, Mid_lse,
    stride_kv: tl.int64, kv_scale,
    sm_scale: tl.constexpr, QK_DIM: tl.constexpr, V_DIM: tl.constexpr,
    BLOCK_K: tl.constexpr, NH: tl.constexpr, V_BLOCK_C: tl.constexpr,
    NUM_SPLITS_C: tl.constexpr, SPLIT_LEN_C: tl.constexpr, KV_SEQ_LEN_C: tl.constexpr,
):
    batch_id = tl.program_id(0)
    split_id = tl.program_id(1)
    v_block_id = tl.program_id(2)

    out_idx = batch_id * NUM_SPLITS_C + split_id
    kv_pos = batch_id * KV_SEQ_LEN_C + split_id * SPLIT_LEN_C
    tl.multiple_of(kv_pos, 128)

    h_offs = tl.arange(0, NH)
    n_offs = tl.arange(0, SPLIT_LEN_C)
    v_start = v_block_id * V_BLOCK_C
    v_offs = v_start + tl.arange(0, V_BLOCK_C)
    v_mask = v_offs < V_DIM
    q_base = batch_id * NH * QK_DIM
    tl.multiple_of(q_base, 128)
    tl.multiple_of(stride_kv, 32)
    tl.assume(stride_kv > 0)

    acc = tl.zeros([NH, V_BLOCK_C], dtype=tl.float32)
    m_i = tl.full([NH], float("-inf"), dtype=tl.float32)
    l_i = tl.zeros([NH], dtype=tl.float32)
    score_scale = sm_scale * kv_scale

    scores = tl.zeros([NH, SPLIT_LEN_C], dtype=tl.float32)
    for k_start in range(0, QK_DIM, BLOCK_K):
        k_offs = tl.arange(0, BLOCK_K)
        k_mask = (k_start + k_offs) < QK_DIM
        q_tile = tl.load(Q_ptr + q_base + h_offs[:, None] * QK_DIM + k_start + k_offs[None, :], mask=k_mask[None, :], other=0.0)
        k_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + k_start + k_offs[None, :], mask=k_mask[None, :], other=0.0)
        scores += tl.dot(q_tile.to(k_tile.dtype), tl.trans(k_tile))

    scores *= score_scale
    row_max = tl.max(scores, axis=1)
    m_new = tl.maximum(m_i, row_max)
    alpha = tl.exp(m_i - m_new)
    l_i = l_i * alpha
    exp_scores = tl.exp(scores - m_new[:, None])
    l_i += tl.sum(exp_scores, axis=1)
    acc = acc * alpha[:, None]

    v_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + v_offs[None, :], mask=v_mask[None, :], other=0.0)
    acc += tl.dot(exp_scores.to(v_tile.dtype), v_tile)
    m_i = m_new

    acc = (acc * kv_scale) / l_i[:, None]
    lse_vals = m_i + tl.log(l_i)
    tl.store(Mid_O + out_idx * NH * V_DIM + h_offs[:, None] * V_DIM + v_offs[None, :], acc.to(tl.bfloat16), mask=v_mask[None, :])
    tl.store(Mid_lse + out_idx * NH + h_offs, lse_vals)


@triton.jit
def _reduce_splitk_serial_s3(
    Mid_O, Mid_lse, O_ptr,
    V_DIM: tl.constexpr, NH: tl.constexpr, BLOCK_V: tl.constexpr, NUM_SPLITS_C: tl.constexpr,
):
    batch_id = tl.program_id(0)
    head_id = 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

    m_final = tl.full([], float("-inf"), dtype=tl.float32)
    l_final = tl.zeros([], dtype=tl.float32)
    acc = tl.zeros([BLOCK_V], dtype=tl.float32)

    for s in range(NUM_SPLITS_C):
        idx = batch_id * NUM_SPLITS_C + s
        lse = tl.load(Mid_lse + idx * NH + head_id)
        m_new = tl.maximum(m_final, lse)
        alpha = tl.exp(m_final - m_new)
        beta = tl.exp(lse - m_new)
        partial = tl.load(Mid_O + idx * NH * V_DIM + head_id * V_DIM + v_offs, mask=v_mask, other=0.0).to(tl.float32)
        acc = acc * alpha + beta * partial
        l_final = l_final * alpha + beta
        m_final = m_new

    acc = acc / l_final
    out_base = batch_id * NH * V_DIM + head_id * V_DIM
    tl.store(O_ptr + out_base + v_offs, acc.to(tl.bfloat16), mask=v_mask)


# ============================================================
# S5 triton kernels (batch=64, kv=1024)
# splits=4, SPLIT_LEN=256, no V-tiling (full 512), serial reduce w2
# ============================================================

S5_NUM_SPLITS = 4
S5_SPLIT_LEN = 256
S5_KV_SEQ_LEN = 1024


@triton.jit
def _flash_decode_fp8_s5_no_vtile(
    Q_ptr, KV_ptr, Mid_O, Mid_lse,
    stride_kv: tl.int64, kv_scale,
    sm_scale: tl.constexpr, QK_DIM: tl.constexpr, V_DIM: tl.constexpr,
    BLOCK_K: tl.constexpr, NH: tl.constexpr,
    NUM_SPLITS_C: tl.constexpr, SPLIT_LEN_C: tl.constexpr, KV_SEQ_LEN_C: tl.constexpr,
):
    batch_id = tl.program_id(0)
    split_id = tl.program_id(1)

    out_idx = batch_id * NUM_SPLITS_C + split_id
    kv_pos = batch_id * KV_SEQ_LEN_C + split_id * SPLIT_LEN_C
    tl.multiple_of(kv_pos, 128)

    h_offs = tl.arange(0, NH)
    n_offs = tl.arange(0, SPLIT_LEN_C)
    v_offs = tl.arange(0, V_DIM)
    q_base = batch_id * NH * QK_DIM
    tl.multiple_of(q_base, 128)
    tl.multiple_of(stride_kv, 32)
    tl.assume(stride_kv > 0)

    acc = tl.zeros([NH, V_DIM], dtype=tl.float32)
    m_i = tl.full([NH], float("-inf"), dtype=tl.float32)
    l_i = tl.zeros([NH], dtype=tl.float32)
    score_scale = sm_scale * kv_scale

    scores = tl.zeros([NH, SPLIT_LEN_C], dtype=tl.float32)
    for k_start in range(0, QK_DIM, BLOCK_K):
        k_offs = tl.arange(0, BLOCK_K)
        k_mask = (k_start + k_offs) < QK_DIM
        q_tile = tl.load(Q_ptr + q_base + h_offs[:, None] * QK_DIM + k_start + k_offs[None, :], mask=k_mask[None, :], other=0.0)
        k_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + k_start + k_offs[None, :], mask=k_mask[None, :], other=0.0)
        scores += tl.dot(q_tile.to(k_tile.dtype), tl.trans(k_tile))

    scores *= score_scale
    row_max = tl.max(scores, axis=1)
    m_new = tl.maximum(m_i, row_max)
    alpha = tl.exp(m_i - m_new)
    l_i = l_i * alpha
    exp_scores = tl.exp(scores - m_new[:, None])
    l_i += tl.sum(exp_scores, axis=1)
    acc = acc * alpha[:, None]

    v_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + v_offs[None, :])
    acc += tl.dot(exp_scores.to(v_tile.dtype), v_tile)
    m_i = m_new

    acc = (acc * kv_scale) / l_i[:, None]
    lse_vals = m_i + tl.log(l_i)
    tl.store(Mid_O + out_idx * NH * V_DIM + h_offs[:, None] * V_DIM + v_offs[None, :], acc.to(tl.bfloat16))
    tl.store(Mid_lse + out_idx * NH + h_offs, lse_vals)


@triton.jit
def _reduce_splitk_serial_s5(
    Mid_O, Mid_lse, O_ptr,
    V_DIM: tl.constexpr, NH: tl.constexpr, BLOCK_V: tl.constexpr, NUM_SPLITS_C: tl.constexpr,
):
    batch_id = tl.program_id(0)
    head_id = 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

    m_final = tl.full([], float("-inf"), dtype=tl.float32)
    l_final = tl.zeros([], dtype=tl.float32)
    acc = tl.zeros([BLOCK_V], dtype=tl.float32)

    for s in range(NUM_SPLITS_C):
        idx = batch_id * NUM_SPLITS_C + s
        lse = tl.load(Mid_lse + idx * NH + head_id)
        m_new = tl.maximum(m_final, lse)
        alpha = tl.exp(m_final - m_new)
        beta = tl.exp(lse - m_new)
        partial = tl.load(Mid_O + idx * NH * V_DIM + head_id * V_DIM + v_offs, mask=v_mask, other=0.0).to(tl.float32)
        acc = acc * alpha + beta * partial
        l_final = l_final * alpha + beta
        m_final = m_new

    acc = acc / l_final
    out_base = batch_id * NH * V_DIM + head_id * V_DIM
    tl.store(O_ptr + out_base + v_offs, acc.to(tl.bfloat16), mask=v_mask)


# ============================================================
# Aiter persistent fp8 path (v143) for S4, S6, S7, S8
# ============================================================

def _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, persistent_splits, fast_mode, kv_gran=64):
    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]


# ============================================================
# Triton cache helpers
# ============================================================

def _ensure_triton_cache(shape_key, batch_size, total_q, num_splits, kv_scale_tensor):
    key = (shape_key, batch_size, total_q)
    if key in _cache:
        return _cache[key]

    _cache[key] = {
        "mid_o": torch.empty((batch_size * num_splits, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
        "mid_lse": torch.empty((batch_size * num_splits, NUM_HEADS), dtype=torch.float32, device="cuda"),
        "o": torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
        "kv_scale_val": kv_scale_tensor.item(),
    }
    return _cache[key]


# ============================================================
# Unified 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

    # ---- S1: batch=4, kv=1024 -> triton s1_v45 ----
    if batch_size <= 4 and kv_seq_len <= 1024:
        kv_fp8, kv_scale = kv_data["fp8"]
        kv_flat = kv_fp8.view(total_kv, QK_HEAD_DIM)
        c = _ensure_triton_cache("s1_v45", batch_size, total_q, S1_NUM_SPLITS, kv_scale)

        grid1 = (batch_size, S1_NUM_SPLITS, triton.cdiv(V_HEAD_DIM, S1_V_BLOCK))
        _flash_decode_fp8_s1_exact_vtile[grid1](
            q, kv_flat, c["mid_o"], c["mid_lse"],
            QK_HEAD_DIM, c["kv_scale_val"], SM_SCALE,
            QK_HEAD_DIM, V_HEAD_DIM, 256, NUM_HEADS,
            S1_V_BLOCK, S1_NUM_SPLITS, S1_SPLIT_LEN, S1_KV_SEQ_LEN,
            num_warps=8, num_stages=3,
        )

        REDUCE_BLOCK_V = 128
        n_v_blocks = triton.cdiv(V_HEAD_DIM, REDUCE_BLOCK_V)
        _reduce_splitk_parallel_s1[(batch_size, NUM_HEADS, n_v_blocks)](
            c["mid_o"], c["mid_lse"], c["o"],
            S1_NUM_SPLITS, V_HEAD_DIM, NUM_HEADS, REDUCE_BLOCK_V,
            num_warps=1,
        )
        return c["o"]

    # ---- S2: batch=4, kv=8192 -> triton s2_v57 ----
    if batch_size <= 4 and kv_seq_len > 1024:
        kv_fp8, kv_scale = kv_data["fp8"]
        kv_flat = kv_fp8.view(total_kv, QK_HEAD_DIM)
        c = _ensure_triton_cache("s2_v57", batch_size, total_q, S2_NUM_SPLITS, kv_scale)

        grid1 = (batch_size, S2_NUM_SPLITS, triton.cdiv(V_HEAD_DIM, S2_V_BLOCK))
        _flash_decode_fp8_s2_exact_vtile[grid1](
            q, kv_flat, c["mid_o"], c["mid_lse"],
            QK_HEAD_DIM, c["kv_scale_val"], SM_SCALE,
            QK_HEAD_DIM, V_HEAD_DIM, 256, 256, NUM_HEADS,
            S2_V_BLOCK, S2_NUM_SPLITS, S2_SPLIT_LEN, S2_KV_SEQ_LEN,
            num_warps=4, num_stages=3,
        )

        REDUCE_BLOCK_V = 256
        n_v_blocks = triton.cdiv(V_HEAD_DIM, REDUCE_BLOCK_V)
        _reduce_splitk_parallel_s2[(batch_size, NUM_HEADS, n_v_blocks)](
            c["mid_o"], c["mid_lse"], c["o"],
            S2_NUM_SPLITS, V_HEAD_DIM, NUM_HEADS, REDUCE_BLOCK_V,
            num_warps=2,
        )
        return c["o"]

    # ---- S3: batch=32, kv=1024 -> triton s3_v37 ----
    if batch_size <= 32 and kv_seq_len <= 1024:
        kv_fp8, kv_scale = kv_data["fp8"]
        kv_flat = kv_fp8.view(total_kv, QK_HEAD_DIM)
        c = _ensure_triton_cache("s3_v37", batch_size, total_q, S3_NUM_SPLITS, kv_scale)

        grid1 = (batch_size, S3_NUM_SPLITS, triton.cdiv(V_HEAD_DIM, S3_V_BLOCK))
        _flash_decode_fp8_s3_exact_vtile[grid1](
            q, kv_flat, c["mid_o"], c["mid_lse"],
            QK_HEAD_DIM, c["kv_scale_val"], SM_SCALE,
            QK_HEAD_DIM, V_HEAD_DIM, 256, NUM_HEADS,
            S3_V_BLOCK, S3_NUM_SPLITS, S3_SPLIT_LEN, S3_KV_SEQ_LEN,
            num_warps=8, num_stages=2,
        )

        REDUCE_BLOCK_V = 256
        n_v_blocks = triton.cdiv(V_HEAD_DIM, REDUCE_BLOCK_V)
        _reduce_splitk_serial_s3[(batch_size, NUM_HEADS, n_v_blocks)](
            c["mid_o"], c["mid_lse"], c["o"],
            V_HEAD_DIM, NUM_HEADS, REDUCE_BLOCK_V, S3_NUM_SPLITS,
            num_warps=1,
        )
        return c["o"]

    # ---- S5: batch=64, kv=1024 -> triton s5_v15 ----
    if batch_size <= 64 and kv_seq_len <= 1024:
        kv_fp8, kv_scale = kv_data["fp8"]
        kv_flat = kv_fp8.view(total_kv, QK_HEAD_DIM)
        c = _ensure_triton_cache("s5_v15", batch_size, total_q, S5_NUM_SPLITS, kv_scale)

        _flash_decode_fp8_s5_no_vtile[(batch_size, S5_NUM_SPLITS)](
            q, kv_flat, c["mid_o"], c["mid_lse"],
            QK_HEAD_DIM, c["kv_scale_val"], SM_SCALE,
            QK_HEAD_DIM, V_HEAD_DIM, 128, NUM_HEADS,
            S5_NUM_SPLITS, S5_SPLIT_LEN, S5_KV_SEQ_LEN,
            num_warps=8, num_stages=3,
        )

        reduce_block_v = 512
        n_v_blocks = triton.cdiv(V_HEAD_DIM, reduce_block_v)
        _reduce_splitk_serial_s5[(batch_size, NUM_HEADS, n_v_blocks)](
            c["mid_o"], c["mid_lse"], c["o"],
            V_HEAD_DIM, NUM_HEADS, reduce_block_v, S5_NUM_SPLITS,
            num_warps=2,
        )
        return c["o"]

    # ---- S4, S6, S7, S8: aiter persistent fp8 (v143) ----
    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)

    if total_kv >= 1000000:
        # s8 (256, 8192)
        splits, fast_mode = 4, False
    elif total_kv >= 300000:
        # s6 (64, 8192)
        splits, fast_mode = 4, False
    elif batch_size >= 256:
        # s7 (256, 1024)
        splits, fast_mode = 4, False
    else:
        # s4 (32, 8192)
        splits, fast_mode = 32, True

    kv_gran = 64
    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 · 652 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