Skip to content
KernelIndex
Search⌘K

submission 588924

divc13 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

sub120_hybrid.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-588924?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.2µs
#464 of 766
2026-03-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:df0d5b2279f1197ba978d9202b82340d5ad49219489c2d5ec7d5b884e708bb39
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-kdef _flash_splitk(
tile-m = 16BLOCK_M = 16 if total_m <= 16 else (32 if total_m <= 32 else 64)

Kernel source

sub120_hybrid.py358 lines
"""
sub120: Hybrid kernel — dispatches between batched bmm and Triton flash attention.

Key insight from benchmarks:
- bmm is VERY fast for small kv_len (bs=4,kv=1024: 23.6µs vs Triton ~35µs)
- bmm is VERY slow for large kv_len (bs=128,kv=8192,qseq=4: 976µs vs Triton ~574µs)

Strategy:
- If kv_len * qseq <= threshold: use batched bmm (avoids kernel launch overhead)
- Else: use Triton flash attention (avoids materializing full score matrix)

Also uses fp8 KV for both paths where possible.
"""

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
SM_SCALE_LOG2E = SM_SCALE * LOG2E


# ==================== Triton Flash Attention (from sub111) ====================

@triton.jit
def _flash_fused(
    Q_ptr, KV_ptr, O_ptr,
    qo_indptr_ptr, kv_indptr_ptr,
    sm_scale_log2e,
    stride_q0, stride_q1,
    stride_kv0,
    stride_o0, stride_o1,
    num_heads: 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)

    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

    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(0, kv_len, BLOCK_KV):
        kv_range = tl.arange(0, BLOCK_KV)
        kv_valid = (kv_off + kv_range) < kv_len
        kv_base = (kv_start + kv_off + kv_range) * stride_kv0

        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'))

        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

        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)

    result = acc / l_i[:, None]
    o_base = qi_global * stride_o0 + hi * stride_o1
    v_range = tl.arange(0, V_DIM)
    tl.store(O_ptr + o_base[:, None] + v_range[None, :],
             result.to(tl.bfloat16), mask=m_mask[:, None])


@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,
    stride_kv0,
    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_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)

    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) * stride_kv0

        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'))

        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

        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)

    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

    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):
        m_s = tl.load(Max_ptr + (base + s) * BLOCK_M + m_range, mask=m_mask, other=float('-inf'))
        global_max = tl.maximum(global_max, m_s)

    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]

    result = total_acc / total_l[:, None]
    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])


# ==================== BMM Path ====================

def _bmm_attention(q, kv_bf16, qo_indptr, kv_indptr, config):
    num_heads = config["num_heads"]
    v_head_dim = config["v_head_dim"]
    batch_size = config["batch_size"]
    q_seq_len = config["q_seq_len"]
    kv_seq_len = config["kv_seq_len"]
    total_q = q.shape[0]

    q_batched = q.view(batch_size, q_seq_len, num_heads, 576).reshape(batch_size, q_seq_len * num_heads, 576)
    kv_batched = kv_bf16.view(batch_size, kv_seq_len, 576)

    scores = torch.bmm(q_batched, kv_batched.transpose(1, 2))
    scores.mul_(SM_SCALE)
    scores = F.softmax(scores, dim=-1)

    v_batched = kv_batched[:, :, :v_head_dim]
    output = torch.bmm(scores.to(v_batched.dtype), v_batched)

    return output.view(batch_size, q_seq_len, num_heads, v_head_dim).reshape(total_q, num_heads, v_head_dim).to(torch.bfloat16)


# ==================== Triton Path ====================

def _triton_flash_attention(q, kv_flat, qo_indptr, kv_indptr, config):
    num_heads = config["num_heads"]
    v_head_dim = config["v_head_dim"]
    q_seq_len = config["q_seq_len"]
    batch_size = config["batch_size"]

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

    total_m = q_seq_len * num_heads
    BLOCK_M = 16 if total_m <= 16 else (32 if total_m <= 32 else 64)
    BLOCK_KV = 64
    D_TILE = 64

    num_m_groups = (total_m + BLOCK_M - 1) // BLOCK_M
    total_programs_base = batch_size * num_m_groups

    if total_programs_base >= 128:
        grid = (batch_size, num_m_groups)
        _flash_fused[grid](
            q, kv_flat, o, qo_indptr, kv_indptr,
            SM_SCALE_LOG2E,
            q.stride(0), q.stride(1), kv_flat.stride(0),
            o.stride(0), o.stride(1),
            num_heads=num_heads, BLOCK_M=BLOCK_M, BLOCK_KV=BLOCK_KV,
            D_TILE=D_TILE, V_DIM=512, HEAD_DIM=576,
        )
    else:
        num_splits = max(1, min(32, 512 // max(1, total_programs_base)))
        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")

        _flash_splitk[(batch_size, num_m_groups, num_splits)](
            q, kv_flat, acc_partial, max_partial, sum_partial,
            qo_indptr, kv_indptr, SM_SCALE_LOG2E,
            q.stride(0), q.stride(1), kv_flat.stride(0),
            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,
        )
        _reduce_splitk[(batch_size, num_m_groups)](
            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


# ==================== Dispatch ====================

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

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

    # Heuristic: bmm is better when the score matrix is small
    # score_matrix_size = batch_size * q_seq_len * num_heads * kv_seq_len
    # bmm materializes the full score matrix in memory
    # Flash attention doesn't, so it wins for large score matrices
    score_size = q_seq_len * kv_seq_len

    if score_size <= 4096:  # e.g., qseq=1, kv≤4096 or qseq=4, kv≤1024
        # Use batched bmm — faster for small problems
        kv_bf16 = kv_data["bf16"]
        return _bmm_attention(q, kv_bf16, qo_indptr, kv_indptr, config)
    else:
        # Use Triton flash attention — better for large score matrices
        kv_bf16 = kv_data["bf16"]
        kv_flat = kv_bf16.view(-1, 576)
        return _triton_flash_attention(q, kv_flat, qo_indptr, kv_indptr, config)
scrolls · 358 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 586148.

"""
- sub75: Triton flash attention with split-K for CU utilization.
+ sub120: Hybrid kernel — dispatches between batched bmm and Triton flash attention.
- 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)
+ Key insight from benchmarks:
+ - bmm is VERY fast for small kv_len (bs=4,kv=1024: 23.6µs vs Triton ~35µs)
+ - bmm is VERY slow for large kv_len (bs=128,kv=8192,qseq=4: 976µs vs Triton ~574µs)
- 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.
+ Strategy:
+ - If kv_len * qseq <= threshold: use batched bmm (avoids kernel launch overhead)
+ - Else: use Triton flash attention (avoids materializing full score matrix)
+
+ Also uses fp8 KV for both paths where possible.
"""
import torch
⋯ 4 unchanged lines
SM_SCALE = 1.0 / (576 ** 0.5)
LOG2E = 1.4426950408889634
+ SM_SCALE_LOG2E = SM_SCALE * LOG2E
+ # ==================== Triton Flash Attention (from sub111) ====================
+
@triton.jit
+ def _flash_fused(
+ Q_ptr, KV_ptr, O_ptr,
+ qo_indptr_ptr, kv_indptr_ptr,
+ sm_scale_log2e,
+ stride_q0, stride_q1,
+ stride_kv0,
+ stride_o0, stride_o1,
+ num_heads: 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)
+
+ 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
+
+ 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(0, kv_len, BLOCK_KV):
+ kv_range = tl.arange(0, BLOCK_KV)
+ kv_valid = (kv_off + kv_range) < kv_len
+ kv_base = (kv_start + kv_off + kv_range) * stride_kv0
+
+ 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'))
+
+ 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
+
+ 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)
+
+ result = acc / l_i[:, None]
+ o_base = qi_global * stride_o0 + hi * stride_o1
+ v_range = tl.arange(0, V_DIM)
+ tl.store(O_ptr + o_base[:, None] + v_range[None, :],
+ result.to(tl.bfloat16), mask=m_mask[:, None])
+
+
+ @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,
+ stride_kv0,
num_heads: tl.constexpr,
num_splits: tl.constexpr,
num_m_groups: tl.constexpr,
⋯ 25 unchanged lines
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)
⋯ 1 unchanged lines
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
+ kv_base = (kv_start + kv_off + kv_range) * stride_kv0
- # 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)
⋯ 10 unchanged lines
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)
⋯ 2 unchanged lines
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, :],
⋯ 1 unchanged lines
).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
⋯ 18 unchanged lines
):
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
⋯ 3 unchanged lines
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'))
+ m_s = tl.load(Max_ptr + (base + s) * 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)
+ 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)
+ 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
+ # ==================== BMM Path ====================
+ def _bmm_attention(q, kv_bf16, qo_indptr, kv_indptr, config):
num_heads = config["num_heads"]
v_head_dim = config["v_head_dim"]
- q_seq_len = config["q_seq_len"]
batch_size = config["batch_size"]
+ q_seq_len = config["q_seq_len"]
+ kv_seq_len = config["kv_seq_len"]
+ total_q = q.shape[0]
- # Use bf16 KV (faster than fp8 in Triton on AMD)
- kv_bf16 = kv_data["bf16"]
- kv_flat = kv_bf16.view(-1, 576)
+ q_batched = q.view(batch_size, q_seq_len, num_heads, 576).reshape(batch_size, q_seq_len * num_heads, 576)
+ kv_batched = kv_bf16.view(batch_size, kv_seq_len, 576)
+ scores = torch.bmm(q_batched, kv_batched.transpose(1, 2))
+ scores.mul_(SM_SCALE)
+ scores = F.softmax(scores, dim=-1)
+
+ v_batched = kv_batched[:, :, :v_head_dim]
+ output = torch.bmm(scores.to(v_batched.dtype), v_batched)
+
+ return output.view(batch_size, q_seq_len, num_heads, v_head_dim).reshape(total_q, num_heads, v_head_dim).to(torch.bfloat16)
+
+
+ # ==================== Triton Path ====================
+
+ def _triton_flash_attention(q, kv_flat, qo_indptr, kv_indptr, config):
+ num_heads = config["num_heads"]
+ v_head_dim = config["v_head_dim"]
+ q_seq_len = config["q_seq_len"]
+ batch_size = config["batch_size"]
+
total_q = q.shape[0]
o = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device="cuda")
- BLOCK_M = 16
+ total_m = q_seq_len * num_heads
+ BLOCK_M = 16 if total_m <= 16 else (32 if total_m <= 32 else 64)
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")
+ if total_programs_base >= 128:
+ grid = (batch_size, num_m_groups)
+ _flash_fused[grid](
+ q, kv_flat, o, qo_indptr, kv_indptr,
+ SM_SCALE_LOG2E,
+ q.stride(0), q.stride(1), kv_flat.stride(0),
+ o.stride(0), o.stride(1),
+ num_heads=num_heads, BLOCK_M=BLOCK_M, BLOCK_KV=BLOCK_KV,
+ D_TILE=D_TILE, V_DIM=512, HEAD_DIM=576,
+ )
+ else:
+ num_splits = max(1, min(32, 512 // max(1, total_programs_base)))
+ 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,
- )
+ _flash_splitk[(batch_size, num_m_groups, num_splits)](
+ q, kv_flat, acc_partial, max_partial, sum_partial,
+ qo_indptr, kv_indptr, SM_SCALE_LOG2E,
+ q.stride(0), q.stride(1), kv_flat.stride(0),
+ 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,
+ )
+ _reduce_splitk[(batch_size, num_m_groups)](
+ 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,
+ )
- # 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
+
+
+ # ==================== Dispatch ====================
+
+ def custom_kernel(data: input_t) -> output_t:
+ q, kv_data, qo_indptr, kv_indptr, config = data
+
+ kv_seq_len = config["kv_seq_len"]
+ q_seq_len = config["q_seq_len"]
+ batch_size = config["batch_size"]
+ num_heads = config["num_heads"]
+
+ # Heuristic: bmm is better when the score matrix is small
+ # score_matrix_size = batch_size * q_seq_len * num_heads * kv_seq_len
+ # bmm materializes the full score matrix in memory
+ # Flash attention doesn't, so it wins for large score matrices
+ score_size = q_seq_len * kv_seq_len
+
+ if score_size <= 4096: # e.g., qseq=1, kv≤4096 or qseq=4, kv≤1024
+ # Use batched bmm — faster for small problems
+ kv_bf16 = kv_data["bf16"]
+ return _bmm_attention(q, kv_bf16, qo_indptr, kv_indptr, config)
+ else:
+ # Use Triton flash attention — better for large score matrices
+ kv_bf16 = kv_data["bf16"]
+ kv_flat = kv_bf16.view(-1, 576)
+ return _triton_flash_attention(q, kv_flat, qo_indptr, kv_indptr, config)
scrolls · 384 diff lines total

Best evidence level for this revision: reported

JSON