Skip to content
KernelIndex
Search⌘K

submission 586148

divc13 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

sub75_triton_splitk.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-586148?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
102.9µs
#467 of 766
2026-03-18

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:947d94a5963105ca4212c465e88ecb08f29b0175b12c2340a7a5f33b3d8dd700
license declaredunknown
license concludedunknown
authorsdivc13
imported2026-08-15

Techniques

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

mmascores += tl.dot(q_chunk, tl.trans(k_chunk))
split-ksub75: Triton flash attention with split-K for CU utilization.
tile-m = 16BLOCK_M = 16

Kernel source

sub75_triton_splitk.py256 lines
"""
sub75: Triton flash attention with split-K for CU utilization.

Key improvements over sub49/sub52:
1. Split-K: divide KV sequence across programs for better parallelism
2. Adaptive num_splits: more splits for small batches, fewer for large
3. Torch-based reduction: simple partial result combination
4. Uses bf16 KV (faster than fp8 on Triton/AMD per sub52 results)

Grid: (batch_size, num_m_groups, num_splits)
Each program handles BLOCK_M M-rows and kv_len/num_splits KV tokens.
Outputs partial (acc, max, sumexp) to global memory.
Reduction combines partials with online softmax correction.
"""

import torch
import torch.nn.functional as F
import triton
import triton.language as tl
from task import input_t, output_t

SM_SCALE = 1.0 / (576 ** 0.5)
LOG2E = 1.4426950408889634


@triton.jit
def _flash_splitk(
    Q_ptr, KV_ptr,
    Acc_ptr, Max_ptr, Sum_ptr,
    qo_indptr_ptr, kv_indptr_ptr,
    sm_scale_log2e,
    stride_q0, stride_q1,
    num_heads: tl.constexpr,
    num_splits: tl.constexpr,
    num_m_groups: tl.constexpr,
    BLOCK_M: tl.constexpr,
    BLOCK_KV: tl.constexpr,
    D_TILE: tl.constexpr,
    V_DIM: tl.constexpr,
    HEAD_DIM: tl.constexpr,
):
    batch = tl.program_id(0)
    m_group = tl.program_id(1)
    split = tl.program_id(2)

    kv_start = tl.load(kv_indptr_ptr + batch)
    kv_end = tl.load(kv_indptr_ptr + batch + 1)
    kv_len = kv_end - kv_start
    q_start = tl.load(qo_indptr_ptr + batch)
    q_end = tl.load(qo_indptr_ptr + batch + 1)
    q_len = q_end - q_start

    total_m = q_len * num_heads
    m_start = m_group * BLOCK_M
    m_range = tl.arange(0, BLOCK_M)
    m_idx = m_start + m_range
    m_mask = m_idx < total_m

    qi_local = m_idx // num_heads
    hi = m_idx % num_heads
    qi_global = q_start + qi_local
    q_base = qi_global * stride_q0 + hi * stride_q1

    # KV range for this split
    kv_per_split = (kv_len + num_splits - 1) // num_splits
    split_kv_start = split * kv_per_split
    split_kv_end = tl.minimum(split_kv_start + kv_per_split, kv_len)

    # Online softmax state
    m_i = tl.full([BLOCK_M], float('-inf'), dtype=tl.float32)
    l_i = tl.zeros([BLOCK_M], dtype=tl.float32)
    acc = tl.zeros([BLOCK_M, V_DIM], dtype=tl.float32)

    for kv_off in range(split_kv_start, split_kv_end, BLOCK_KV):
        kv_range = tl.arange(0, BLOCK_KV)
        kv_valid = (kv_off + kv_range) < split_kv_end
        kv_base = (kv_start + kv_off + kv_range) * HEAD_DIM

        # QK^T tiled by D_TILE
        scores = tl.zeros([BLOCK_M, BLOCK_KV], dtype=tl.float32)
        for d_off in tl.static_range(0, HEAD_DIM, D_TILE):
            d_range = tl.arange(0, D_TILE)
            q_chunk = tl.load(
                Q_ptr + q_base[:, None] + d_off + d_range[None, :],
                mask=m_mask[:, None], other=0.0
            ).to(tl.bfloat16)
            k_chunk = tl.load(
                KV_ptr + kv_base[:, None] + d_off + d_range[None, :],
                mask=kv_valid[:, None], other=0.0
            ).to(tl.bfloat16)
            scores += tl.dot(q_chunk, tl.trans(k_chunk))

        scores *= sm_scale_log2e
        scores = tl.where(kv_valid[None, :], scores, float('-inf'))

        # Online softmax
        m_ij = tl.max(scores, axis=1)
        new_m = tl.maximum(m_i, m_ij)
        alpha = tl.math.exp2(m_i - new_m)
        p = tl.math.exp2(scores - new_m[:, None])
        l_i = l_i * alpha + tl.sum(p, axis=1)
        acc = acc * alpha[:, None]
        m_i = new_m

        # OV: V = first 512 dims of KV
        v_range = tl.arange(0, V_DIM)
        v_block = tl.load(
            KV_ptr + kv_base[:, None] + v_range[None, :],
            mask=kv_valid[:, None], other=0.0
        ).to(tl.bfloat16)
        acc += tl.dot(p.to(tl.bfloat16), v_block)

    # Store partial results
    # Layout: flat index = (batch * num_m_groups + m_group) * num_splits + split
    flat_idx = (batch * num_m_groups + m_group) * num_splits + split
    acc_base = flat_idx * BLOCK_M * V_DIM
    ml_base = flat_idx * BLOCK_M

    v_range = tl.arange(0, V_DIM)
    tl.store(Acc_ptr + acc_base + m_range[:, None] * V_DIM + v_range[None, :],
             acc, mask=m_mask[:, None])
    tl.store(Max_ptr + ml_base + m_range, m_i, mask=m_mask)
    tl.store(Sum_ptr + ml_base + m_range, l_i, mask=m_mask)


@triton.jit
def _reduce_splitk(
    Acc_ptr, Max_ptr, Sum_ptr, O_ptr,
    qo_indptr_ptr,
    stride_o0, stride_o1,
    num_heads: tl.constexpr,
    num_splits: tl.constexpr,
    num_m_groups: tl.constexpr,
    BLOCK_M: tl.constexpr,
    V_DIM: tl.constexpr,
):
    batch = tl.program_id(0)
    m_group = tl.program_id(1)

    q_start = tl.load(qo_indptr_ptr + batch)
    q_end = tl.load(qo_indptr_ptr + batch + 1)
    q_len = q_end - q_start
    total_m = q_len * num_heads

    m_start = m_group * BLOCK_M
    m_range = tl.arange(0, BLOCK_M)
    m_idx = m_start + m_range
    m_mask = m_idx < total_m

    qi_local = m_idx // num_heads
    hi = m_idx % num_heads
    qi_global = q_start + qi_local

    # Find global max across splits
    base = (batch * num_m_groups + m_group) * num_splits
    global_max = tl.full([BLOCK_M], float('-inf'), dtype=tl.float32)
    for s in range(num_splits):
        flat_idx = base + s
        m_s = tl.load(Max_ptr + flat_idx * BLOCK_M + m_range,
                       mask=m_mask, other=float('-inf'))
        global_max = tl.maximum(global_max, m_s)

    # Combine partials
    v_range = tl.arange(0, V_DIM)
    total_acc = tl.zeros([BLOCK_M, V_DIM], dtype=tl.float32)
    total_l = tl.zeros([BLOCK_M], dtype=tl.float32)

    for s in range(num_splits):
        flat_idx = base + s
        m_s = tl.load(Max_ptr + flat_idx * BLOCK_M + m_range,
                       mask=m_mask, other=float('-inf'))
        l_s = tl.load(Sum_ptr + flat_idx * BLOCK_M + m_range,
                       mask=m_mask, other=0.0)
        alpha = tl.math.exp2(m_s - global_max)
        total_l += l_s * alpha

        acc_base = flat_idx * BLOCK_M * V_DIM
        acc_s = tl.load(
            Acc_ptr + acc_base + m_range[:, None] * V_DIM + v_range[None, :],
            mask=m_mask[:, None], other=0.0)
        total_acc += acc_s * alpha[:, None]

    # Normalize
    result = total_acc / total_l[:, None]

    # Store final output
    o_base = qi_global * stride_o0 + hi * stride_o1
    tl.store(O_ptr + o_base[:, None] + v_range[None, :],
             result.to(tl.bfloat16), mask=m_mask[:, None])


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

    num_heads = config["num_heads"]
    v_head_dim = config["v_head_dim"]
    q_seq_len = config["q_seq_len"]
    batch_size = config["batch_size"]

    # Use bf16 KV (faster than fp8 in Triton on AMD)
    kv_bf16 = kv_data["bf16"]
    kv_flat = kv_bf16.view(-1, 576)

    total_q = q.shape[0]
    o = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device="cuda")

    BLOCK_M = 16
    BLOCK_KV = 64
    D_TILE = 64

    total_m = q_seq_len * num_heads
    num_m_groups = (total_m + BLOCK_M - 1) // BLOCK_M

    # Adaptive split-K: target ~512 total programs for good CU utilization
    total_programs_base = batch_size * num_m_groups
    num_splits = max(1, min(32, 512 // max(1, total_programs_base)))

    # Allocate partial buffers
    total_partials = batch_size * num_m_groups * num_splits
    acc_partial = torch.empty((total_partials * BLOCK_M, 512), dtype=torch.float32, device="cuda")
    max_partial = torch.empty((total_partials * BLOCK_M,), dtype=torch.float32, device="cuda")
    sum_partial = torch.empty((total_partials * BLOCK_M,), dtype=torch.float32, device="cuda")

    # Launch flash attention kernel
    grid_flash = (batch_size, num_m_groups, num_splits)
    _flash_splitk[grid_flash](
        q, kv_flat,
        acc_partial, max_partial, sum_partial,
        qo_indptr, kv_indptr,
        SM_SCALE * LOG2E,
        q.stride(0), q.stride(1),
        num_heads=num_heads,
        num_splits=num_splits,
        num_m_groups=num_m_groups,
        BLOCK_M=BLOCK_M,
        BLOCK_KV=BLOCK_KV,
        D_TILE=D_TILE,
        V_DIM=512,
        HEAD_DIM=576,
    )

    # Launch reduction kernel
    grid_reduce = (batch_size, num_m_groups)
    _reduce_splitk[grid_reduce](
        acc_partial, max_partial, sum_partial, o,
        qo_indptr,
        o.stride(0), o.stride(1),
        num_heads=num_heads,
        num_splits=num_splits,
        num_m_groups=num_m_groups,
        BLOCK_M=BLOCK_M,
        V_DIM=512,
    )

    return o
scrolls · 256 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 585963.

"""
- sub74: Pure PyTorch bmm attention with bf16 KV.
- No custom kernel, no aiter. Tests batched GEMM approach.
+ sub75: Triton flash attention with split-K for CU utilization.
- Advantages:
- - Single torch.bmm call for QK^T (hipBLAS uses MFMA internally)
- - Single torch.bmm call for OV
- - No per-batch loops, no kernel launch overhead
- - MQA: K/V naturally broadcast via batched matmul
+ Key improvements over sub49/sub52:
+ 1. Split-K: divide KV sequence across programs for better parallelism
+ 2. Adaptive num_splits: more splits for small batches, fewer for large
+ 3. Torch-based reduction: simple partial result combination
+ 4. Uses bf16 KV (faster than fp8 on Triton/AMD per sub52 results)
- Disadvantages:
- - bf16 KV = 2x bandwidth of fp8
- - Materializes full [bs, qseq*nh, kv_seq] scores tensor
+ Grid: (batch_size, num_m_groups, num_splits)
+ Each program handles BLOCK_M M-rows and kv_len/num_splits KV tokens.
+ Outputs partial (acc, max, sumexp) to global memory.
+ Reduction combines partials with online softmax correction.
"""
import torch
import torch.nn.functional as F
+ import triton
+ import triton.language as tl
from task import input_t, output_t
SM_SCALE = 1.0 / (576 ** 0.5)
+ LOG2E = 1.4426950408889634
+ @triton.jit
+ def _flash_splitk(
+ Q_ptr, KV_ptr,
+ Acc_ptr, Max_ptr, Sum_ptr,
+ qo_indptr_ptr, kv_indptr_ptr,
+ sm_scale_log2e,
+ stride_q0, stride_q1,
+ num_heads: tl.constexpr,
+ num_splits: tl.constexpr,
+ num_m_groups: tl.constexpr,
+ BLOCK_M: tl.constexpr,
+ BLOCK_KV: tl.constexpr,
+ D_TILE: tl.constexpr,
+ V_DIM: tl.constexpr,
+ HEAD_DIM: tl.constexpr,
+ ):
+ batch = tl.program_id(0)
+ m_group = tl.program_id(1)
+ split = tl.program_id(2)
+
+ kv_start = tl.load(kv_indptr_ptr + batch)
+ kv_end = tl.load(kv_indptr_ptr + batch + 1)
+ kv_len = kv_end - kv_start
+ q_start = tl.load(qo_indptr_ptr + batch)
+ q_end = tl.load(qo_indptr_ptr + batch + 1)
+ q_len = q_end - q_start
+
+ total_m = q_len * num_heads
+ m_start = m_group * BLOCK_M
+ m_range = tl.arange(0, BLOCK_M)
+ m_idx = m_start + m_range
+ m_mask = m_idx < total_m
+
+ qi_local = m_idx // num_heads
+ hi = m_idx % num_heads
+ qi_global = q_start + qi_local
+ q_base = qi_global * stride_q0 + hi * stride_q1
+
+ # KV range for this split
+ kv_per_split = (kv_len + num_splits - 1) // num_splits
+ split_kv_start = split * kv_per_split
+ split_kv_end = tl.minimum(split_kv_start + kv_per_split, kv_len)
+
+ # Online softmax state
+ m_i = tl.full([BLOCK_M], float('-inf'), dtype=tl.float32)
+ l_i = tl.zeros([BLOCK_M], dtype=tl.float32)
+ acc = tl.zeros([BLOCK_M, V_DIM], dtype=tl.float32)
+
+ for kv_off in range(split_kv_start, split_kv_end, BLOCK_KV):
+ kv_range = tl.arange(0, BLOCK_KV)
+ kv_valid = (kv_off + kv_range) < split_kv_end
+ kv_base = (kv_start + kv_off + kv_range) * HEAD_DIM
+
+ # QK^T tiled by D_TILE
+ scores = tl.zeros([BLOCK_M, BLOCK_KV], dtype=tl.float32)
+ for d_off in tl.static_range(0, HEAD_DIM, D_TILE):
+ d_range = tl.arange(0, D_TILE)
+ q_chunk = tl.load(
+ Q_ptr + q_base[:, None] + d_off + d_range[None, :],
+ mask=m_mask[:, None], other=0.0
+ ).to(tl.bfloat16)
+ k_chunk = tl.load(
+ KV_ptr + kv_base[:, None] + d_off + d_range[None, :],
+ mask=kv_valid[:, None], other=0.0
+ ).to(tl.bfloat16)
+ scores += tl.dot(q_chunk, tl.trans(k_chunk))
+
+ scores *= sm_scale_log2e
+ scores = tl.where(kv_valid[None, :], scores, float('-inf'))
+
+ # Online softmax
+ m_ij = tl.max(scores, axis=1)
+ new_m = tl.maximum(m_i, m_ij)
+ alpha = tl.math.exp2(m_i - new_m)
+ p = tl.math.exp2(scores - new_m[:, None])
+ l_i = l_i * alpha + tl.sum(p, axis=1)
+ acc = acc * alpha[:, None]
+ m_i = new_m
+
+ # OV: V = first 512 dims of KV
+ v_range = tl.arange(0, V_DIM)
+ v_block = tl.load(
+ KV_ptr + kv_base[:, None] + v_range[None, :],
+ mask=kv_valid[:, None], other=0.0
+ ).to(tl.bfloat16)
+ acc += tl.dot(p.to(tl.bfloat16), v_block)
+
+ # Store partial results
+ # Layout: flat index = (batch * num_m_groups + m_group) * num_splits + split
+ flat_idx = (batch * num_m_groups + m_group) * num_splits + split
+ acc_base = flat_idx * BLOCK_M * V_DIM
+ ml_base = flat_idx * BLOCK_M
+
+ v_range = tl.arange(0, V_DIM)
+ tl.store(Acc_ptr + acc_base + m_range[:, None] * V_DIM + v_range[None, :],
+ acc, mask=m_mask[:, None])
+ tl.store(Max_ptr + ml_base + m_range, m_i, mask=m_mask)
+ tl.store(Sum_ptr + ml_base + m_range, l_i, mask=m_mask)
+
+
+ @triton.jit
+ def _reduce_splitk(
+ Acc_ptr, Max_ptr, Sum_ptr, O_ptr,
+ qo_indptr_ptr,
+ stride_o0, stride_o1,
+ num_heads: tl.constexpr,
+ num_splits: tl.constexpr,
+ num_m_groups: tl.constexpr,
+ BLOCK_M: tl.constexpr,
+ V_DIM: tl.constexpr,
+ ):
+ batch = tl.program_id(0)
+ m_group = tl.program_id(1)
+
+ q_start = tl.load(qo_indptr_ptr + batch)
+ q_end = tl.load(qo_indptr_ptr + batch + 1)
+ q_len = q_end - q_start
+ total_m = q_len * num_heads
+
+ m_start = m_group * BLOCK_M
+ m_range = tl.arange(0, BLOCK_M)
+ m_idx = m_start + m_range
+ m_mask = m_idx < total_m
+
+ qi_local = m_idx // num_heads
+ hi = m_idx % num_heads
+ qi_global = q_start + qi_local
+
+ # Find global max across splits
+ base = (batch * num_m_groups + m_group) * num_splits
+ global_max = tl.full([BLOCK_M], float('-inf'), dtype=tl.float32)
+ for s in range(num_splits):
+ flat_idx = base + s
+ m_s = tl.load(Max_ptr + flat_idx * BLOCK_M + m_range,
+ mask=m_mask, other=float('-inf'))
+ global_max = tl.maximum(global_max, m_s)
+
+ # Combine partials
+ v_range = tl.arange(0, V_DIM)
+ total_acc = tl.zeros([BLOCK_M, V_DIM], dtype=tl.float32)
+ total_l = tl.zeros([BLOCK_M], dtype=tl.float32)
+
+ for s in range(num_splits):
+ flat_idx = base + s
+ m_s = tl.load(Max_ptr + flat_idx * BLOCK_M + m_range,
+ mask=m_mask, other=float('-inf'))
+ l_s = tl.load(Sum_ptr + flat_idx * BLOCK_M + m_range,
+ mask=m_mask, other=0.0)
+ alpha = tl.math.exp2(m_s - global_max)
+ total_l += l_s * alpha
+
+ acc_base = flat_idx * BLOCK_M * V_DIM
+ acc_s = tl.load(
+ Acc_ptr + acc_base + m_range[:, None] * V_DIM + v_range[None, :],
+ mask=m_mask[:, None], other=0.0)
+ total_acc += acc_s * alpha[:, None]
+
+ # Normalize
+ result = total_acc / total_l[:, None]
+
+ # Store final output
+ o_base = qi_global * stride_o0 + hi * stride_o1
+ tl.store(O_ptr + o_base[:, None] + v_range[None, :],
+ result.to(tl.bfloat16), mask=m_mask[:, None])
+
+
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
- bs = config["batch_size"]
- nh = config["num_heads"]
- qseq = config["q_seq_len"]
- kv_lora_rank = config["kv_lora_rank"]
- kv_seq = config["kv_seq_len"]
+ num_heads = config["num_heads"]
+ v_head_dim = config["v_head_dim"]
+ q_seq_len = config["q_seq_len"]
+ batch_size = config["batch_size"]
- kv_bf16 = kv_data["bf16"] # [total_kv, 1, 576] bf16
+ # Use bf16 KV (faster than fp8 in Triton on AMD)
+ kv_bf16 = kv_data["bf16"]
+ kv_flat = kv_bf16.view(-1, 576)
- # Reshape for batched matmul (assumes uniform kv_len across batch)
- Q = q.view(bs, qseq, nh, 576).reshape(bs, qseq * nh, 576) # [bs, M, 576]
- K = kv_bf16.view(bs, kv_seq, 576) # [bs, N, 576]
- V = K[:, :, :kv_lora_rank] # [bs, N, 512]
+ total_q = q.shape[0]
+ o = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device="cuda")
- # QK^T: [bs, M, N] via bf16 GEMM (hipBLAS → MFMA)
- scores = torch.bmm(Q, K.transpose(1, 2))
- # Softmax in fp32 for numerical stability
- scores = F.softmax(scores.float() * SM_SCALE, dim=-1)
+ BLOCK_M = 16
+ BLOCK_KV = 64
+ D_TILE = 64
- # OV: [bs, M, 512] via bf16 GEMM
- output = torch.bmm(scores.to(torch.bfloat16), V)
+ total_m = q_seq_len * num_heads
+ num_m_groups = (total_m + BLOCK_M - 1) // BLOCK_M
- return output.view(bs, qseq, nh, 512).reshape(bs * qseq, nh, 512)
+ # Adaptive split-K: target ~512 total programs for good CU utilization
+ total_programs_base = batch_size * num_m_groups
+ num_splits = max(1, min(32, 512 // max(1, total_programs_base)))
+
+ # Allocate partial buffers
+ total_partials = batch_size * num_m_groups * num_splits
+ acc_partial = torch.empty((total_partials * BLOCK_M, 512), dtype=torch.float32, device="cuda")
+ max_partial = torch.empty((total_partials * BLOCK_M,), dtype=torch.float32, device="cuda")
+ sum_partial = torch.empty((total_partials * BLOCK_M,), dtype=torch.float32, device="cuda")
+
+ # Launch flash attention kernel
+ grid_flash = (batch_size, num_m_groups, num_splits)
+ _flash_splitk[grid_flash](
+ q, kv_flat,
+ acc_partial, max_partial, sum_partial,
+ qo_indptr, kv_indptr,
+ SM_SCALE * LOG2E,
+ q.stride(0), q.stride(1),
+ num_heads=num_heads,
+ num_splits=num_splits,
+ num_m_groups=num_m_groups,
+ BLOCK_M=BLOCK_M,
+ BLOCK_KV=BLOCK_KV,
+ D_TILE=D_TILE,
+ V_DIM=512,
+ HEAD_DIM=576,
+ )
+
+ # Launch reduction kernel
+ grid_reduce = (batch_size, num_m_groups)
+ _reduce_splitk[grid_reduce](
+ acc_partial, max_partial, sum_partial, o,
+ qo_indptr,
+ o.stride(0), o.stride(1),
+ num_heads=num_heads,
+ num_splits=num_splits,
+ num_m_groups=num_m_groups,
+ BLOCK_M=BLOCK_M,
+ V_DIM=512,
+ )
+
+ return o
scrolls · 282 diff lines total

Best evidence level for this revision: reported

JSON