Skip to content
KernelIndex
Search⌘K

submission 650700

kosox97741 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

test_v109_blockk128.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-650700?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
80.4µs
#398 of 766
2026-03-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:115486f2d0798c94f4c7369070c151b05aa6e218f55b426884ef177a5ede5fd5
license declaredunknown
license concludedunknown
authorskosox97741
imported2026-08-26

Techniques

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

mmascores += tl.dot(q_tile, tl.trans(k_tile))
num-warps = 8num_warps=8,
online-softmaxm_new = tl.maximum(m_i, row_max)
split-k"""Multi-head split-K flash decode. Grid = (batch, splits)."""
stages = 2num_stages=2,
tile-k = 64Rationale: v102 s8 profile: 9 K-loop iterations per N block with BLOCK_K=64.

Kernel source

test_v109_blockk128.py387 lines
"""
test_v109_blockk128: BLOCK_K 64→128 (reduces K-loop from 9 to 5 iterations)
Base: test_v107_block128.py
Direction: CONTINUING from v108 (attempt 12)
Target: ALL shapes — reduce inner K-loop iterations
Change: BLOCK_K 64→128. QK_DIM=576 / 128 = 4.5 → 5 iterations (last one masked).
        Reduces K-loop overhead by 44%. Need K-dimension masking for indices >= 576.
Rationale: v102 s8 profile: 9 K-loop iterations per N block with BLOCK_K=64.
           LDSBankConflict=24.3%. Fewer K iterations may reduce LDS conflicts.
Scale: INCREMENTAL
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t

# MLA constants
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  # 576
V_HEAD_DIM = KV_LORA_RANK  # 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)

_cache = {}

BLOCK_K = 128  # K dim tile (576 / 128 = 5 iterations, last one masked)


# =============================================================================
# Multi-head flash-decode with tl.dot — fp8 KV path
# (Same kernel, BLOCK_N passed as constexpr — Triton JIT caches per-constexpr)
# =============================================================================
@triton.jit
def _flash_decode_multihead_fp8(
    Q_ptr, KV_ptr, KV_scale, kv_indptr,
    Mid_O,       # (batch * splits, NH, V_DIM) fp32
    Mid_lse,     # (batch * splits, NH) fp32
    stride_kv: tl.int64,
    sm_scale: tl.constexpr,
    QK_DIM: tl.constexpr,   # 576
    V_DIM: tl.constexpr,    # 512
    NUM_SPLITS: tl.constexpr,
    BLOCK_N: tl.constexpr,  # 32 or 64
    BLOCK_K: tl.constexpr,  # 64
    NH: tl.constexpr,       # 16
):
    """Multi-head split-K flash decode. Grid = (batch, splits)."""
    batch_id = tl.program_id(0)
    split_id = tl.program_id(1)

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

    split_len = tl.cdiv(kv_len, NUM_SPLITS)
    my_start = kv_start + split_id * split_len
    my_end = tl.minimum(kv_start + (split_id + 1) * split_len, kv_end)
    actual_len = my_end - my_start

    out_idx = batch_id * NUM_SPLITS + split_id
    h_offs = tl.arange(0, NH)
    v_offs = tl.arange(0, V_DIM)

    if actual_len <= 0:
        tl.store(Mid_lse + out_idx * NH + h_offs, tl.full([NH], float("-inf"), dtype=tl.float32))
        return

    kv_scale_val = tl.load(KV_scale).to(tl.float32)
    combined_scale = kv_scale_val * sm_scale

    q_base = batch_id * NH * QK_DIM

    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)

    for kv_offset in range(0, actual_len, BLOCK_N):
        n_valid = tl.minimum(BLOCK_N, actual_len - kv_offset)
        kv_pos = my_start + kv_offset
        n_offs = tl.arange(0, BLOCK_N)
        n_mask = n_offs < n_valid

        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  # mask for last iteration (576 % 128 = 64)
            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,
            ).to(tl.bfloat16)
            k_tile = tl.load(
                KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + k_start + k_offs[None, :],
                mask=n_mask[:, None] & k_mask[None, :],
                other=0.0,
            ).to(tl.bfloat16)
            scores += tl.dot(q_tile, tl.trans(k_tile))

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

        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=n_mask[:, None],
            other=0.0,
        ).to(tl.bfloat16)
        acc += tl.dot(exp_scores.to(tl.bfloat16), v_tile)

        m_i = m_new

    acc = (acc * kv_scale_val) / 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,
    )
    tl.store(Mid_lse + out_idx * NH + h_offs, lse_vals)


# =============================================================================
# Multi-head flash-decode — bf16 KV path
# =============================================================================
@triton.jit
def _flash_decode_multihead_bf16(
    Q_ptr, KV_ptr, kv_indptr,
    Mid_O, Mid_lse,
    stride_kv: tl.int64,
    sm_scale: tl.constexpr,
    QK_DIM: tl.constexpr,
    V_DIM: tl.constexpr,
    NUM_SPLITS: tl.constexpr,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
    NH: tl.constexpr,
):
    """Multi-head split-K flash decode with bf16 KV. Grid = (batch, splits)."""
    batch_id = tl.program_id(0)
    split_id = tl.program_id(1)

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

    split_len = tl.cdiv(kv_len, NUM_SPLITS)
    my_start = kv_start + split_id * split_len
    my_end = tl.minimum(kv_start + (split_id + 1) * split_len, kv_end)
    actual_len = my_end - my_start

    out_idx = batch_id * NUM_SPLITS + split_id
    h_offs = tl.arange(0, NH)
    v_offs = tl.arange(0, V_DIM)

    if actual_len <= 0:
        tl.store(Mid_lse + out_idx * NH + h_offs, tl.full([NH], float("-inf"), dtype=tl.float32))
        return

    q_base = batch_id * NH * QK_DIM

    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)

    for kv_offset in range(0, actual_len, BLOCK_N):
        n_valid = tl.minimum(BLOCK_N, actual_len - kv_offset)
        kv_pos = my_start + kv_offset
        n_offs = tl.arange(0, BLOCK_N)
        n_mask = n_offs < n_valid

        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  # mask for last iteration (576 % 128 = 64)
            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,
            ).to(tl.bfloat16)
            k_tile = tl.load(
                KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + k_start + k_offs[None, :],
                mask=n_mask[:, None] & k_mask[None, :],
                other=0.0,
            ).to(tl.bfloat16)
            scores += tl.dot(q_tile, tl.trans(k_tile))

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

        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=n_mask[:, None],
            other=0.0,
        ).to(tl.bfloat16)
        acc += tl.dot(exp_scores.to(tl.bfloat16), v_tile)

        m_i = m_new

    acc = acc / 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,
    )
    tl.store(Mid_lse + out_idx * NH + h_offs, lse_vals)


# =============================================================================
# Split-K reduce kernel
# =============================================================================
@triton.jit
def _reduce_splitk(
    Mid_O,      # (batch * splits, NH, V_DIM) fp32
    Mid_lse,    # (batch * splits, NH) fp32
    O_ptr,      # (total_q, NH, V_DIM) bf16
    NUM_SPLITS: tl.constexpr,
    V_DIM: tl.constexpr,
    NH: tl.constexpr,
    BLOCK_V: tl.constexpr,
):
    """Reduce split-K partials. Grid = (batch, NH, cdiv(V_DIM, BLOCK_V))."""
    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):
        idx = batch_id * NUM_SPLITS + s
        lse = tl.load(Mid_lse + idx * NH + head_id)

        is_valid = lse > float("-inf")
        if is_valid:
            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,
            )
            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)


# =============================================================================
# Cache helper
# =============================================================================
def _ensure_cache(batch_size, kv_seq_len, total_q, num_splits):
    key = ("triton_mh_v109", batch_size, kv_seq_len, num_splits)
    if key in _cache:
        return _cache[key]

    _cache[key] = {
        "mid_o": torch.empty(
            (batch_size * num_splits, NUM_HEADS, V_HEAD_DIM),
            dtype=torch.float32, 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",
        ),
    }
    return _cache[key]


# =============================================================================
# Main dispatch — fully custom, no aiter
# =============================================================================
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

    # Per-shape BLOCK_N: 128 for long KV (halves iterations vs v102's 64), 32 for short
    # Profile v102 s8: VGPR=60 (no spill), safe to increase BLOCK_N
    if kv_seq_len >= 8192:
        block_n = 128
    else:
        block_n = 32

    # Adaptive split count — per-shape tuning (combines best of v100 + v101)
    if kv_seq_len <= 1024 and batch_size <= 4:
        # s1 (4, 1024): need more splits for CU utilization
        # min_tokens=64 → max_splits=16 → programs=64 (matches v100)
        min_tokens_per_split = 2 * block_n  # 64
    elif kv_seq_len <= 1024:
        # s3/s5/s7: use min 128 tokens/split (prevents s5 over-splitting)
        min_tokens_per_split = 4 * block_n  # 128
    else:
        # Long KV (s2/s4/s6/s8): larger blocks, allow more splits
        min_tokens_per_split = 4 * block_n  # 256 with block_n=64

    max_splits = max(1, kv_seq_len // min_tokens_per_split)

    target_programs = 2048
    target_splits = max(1, target_programs // batch_size)
    num_splits = max(1, min(target_splits, max_splits))

    # Cap splits
    if num_splits > 64:
        num_splits = 64

    c = _ensure_cache(batch_size, kv_seq_len, total_q, num_splits)

    if batch_size <= 4:
        # bf16 path — simpler, no fp8 scale overhead
        kv_bf16 = kv_data["bf16"]
        kv_flat = kv_bf16.view(total_kv, QK_HEAD_DIM)

        grid1 = (batch_size, num_splits)
        _flash_decode_multihead_bf16[grid1](
            q, kv_flat, kv_indptr,
            c["mid_o"], c["mid_lse"],
            QK_HEAD_DIM,
            SM_SCALE,
            QK_HEAD_DIM, V_HEAD_DIM,
            num_splits,
            block_n, BLOCK_K, NUM_HEADS,
            num_warps=8,
            num_stages=2,
        )
    else:
        # fp8 path — 2x bandwidth savings
        kv_fp8, kv_scale = kv_data["fp8"]
        kv_flat = kv_fp8.view(total_kv, QK_HEAD_DIM)

        grid1 = (batch_size, num_splits)
        _flash_decode_multihead_fp8[grid1](
            q, kv_flat, kv_scale, kv_indptr,
            c["mid_o"], c["mid_lse"],
            QK_HEAD_DIM,
            SM_SCALE,
            QK_HEAD_DIM, V_HEAD_DIM,
            num_splits,
            block_n, BLOCK_K, NUM_HEADS,
            num_warps=8,
            num_stages=2,
        )

    # Stage 2: reduce split-K partials
    REDUCE_BLOCK_V = 128
    n_v_blocks = triton.cdiv(V_HEAD_DIM, REDUCE_BLOCK_V)
    grid2 = (batch_size, NUM_HEADS, n_v_blocks)
    _reduce_splitk[grid2](
        c["mid_o"], c["mid_lse"], c["o"],
        num_splits, V_HEAD_DIM, NUM_HEADS, REDUCE_BLOCK_V,
        num_warps=4,
    )

    return c["o"]
scrolls · 387 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