Skip to content
KernelIndex
Search⌘K

submission 746709

divc13 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v140.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-746709?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
31.9µs
#23 of 766
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:700e94e76d00aefeca7a07512e2afd271f393db130afa6f9f6aa08e7aaffef26
license declaredunknown
license concludedunknown
authorsdivc13
imported2026-08-15

Techniques

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

fp8FP8_DTYPE = tl.float8e4nv
mmaqk = tl.dot(q_nope, k_nope) + tl.dot(q_rope, k_rope)
num-warps = 4num_warps=4, num_stages=2, waves_per_eu=2,
stages = 2num_warps=4, num_stages=2, waves_per_eu=2,
tile-n = 64NUM_KV_SPLITS=num_splits, BLOCK_N=64, BLOCK_H=16,

Kernel source

submission_v140.py267 lines
"""MLA decode Triton v140: v135 + separated LSE tensor (stride=512).

Isolates the separated tensor layout from v138 (which also changed
stage2 nw=2). Keeps stage2 nw=4 to avoid the bs32/kv8k regression.

mid_acc: [batch, heads, splits, 512] stride=512 (power of 2, shift)
mid_lse: [batch, heads, splits]       separate small tensor
"""

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

NUM_HEADS = 16
KV_LORA_RANK = 512
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
FP8_DTYPE = tl.float8e4nv
LOG2E = tl.constexpr(1.4426950408889634)
LN2 = tl.constexpr(0.6931471805599453)

_SHAPE_CONFIGS = {
    (4, 1024):   16,
    (4, 8192):   32,
    (32, 1024):  4,
    (32, 8192):  8,
    (64, 1024):  4,
    (64, 8192):  8,
    (256, 1024): 1,
    (256, 8192): 2,
}


@triton.jit
def _remap_xcd(pid, GRID_SIZE, NUM_XCDS: tl.constexpr = 8):
    pids_per_xcd = (GRID_SIZE + NUM_XCDS - 1) // NUM_XCDS
    tall_xcds = GRID_SIZE % NUM_XCDS
    tall_xcds = NUM_XCDS if tall_xcds == 0 else tall_xcds
    xcd = pid % NUM_XCDS
    local_pid = pid // NUM_XCDS
    if xcd < tall_xcds:
        return xcd * pids_per_xcd + local_pid
    else:
        return tall_xcds * pids_per_xcd + (xcd - tall_xcds) * (pids_per_xcd - 1) + local_pid


@triton.jit
def _fast_rcp(x):
    return tl.inline_asm_elementwise(
        "v_rcp_f32_e32 $0, $1", "=v, v", [x],
        dtype=tl.float32, is_pure=True, pack=1,
    )


@triton.jit
def _stage1_kernel(
    Q, KV, Mid_Acc, Mid_LSE, Out,
    qo_indptr, kv_indptr,
    sm_scale, kv_scale_ptr,
    NUM_KV_SPLITS: tl.constexpr, BLOCK_N: tl.constexpr,
    BLOCK_H: tl.constexpr, BLOCK_C: tl.constexpr, BLOCK_R: tl.constexpr,
    DIRECT_OUT: tl.constexpr, GRID_SIZE: tl.constexpr,
    STRIDE_QT: tl.constexpr, STRIDE_QH: tl.constexpr,
    STRIDE_KVT: tl.constexpr,
    STRIDE_OT: tl.constexpr, STRIDE_OH: tl.constexpr,
):
    raw_pid = tl.program_id(0)
    pid = _remap_xcd(raw_pid, GRID_SIZE)
    cur_batch = pid // NUM_KV_SPLITS
    split_id = pid % NUM_KV_SPLITS
    heads = tl.arange(0, BLOCK_H)
    kv_scale = tl.load(kv_scale_ptr).to(tl.float32)
    score_scale = sm_scale * kv_scale
    q_tok = tl.load(qo_indptr + cur_batch)
    kv_s = tl.load(kv_indptr + cur_batch)
    kv_e = tl.load(kv_indptr + cur_batch + 1)
    kv_len = kv_e - kv_s
    kv_per_split = tl.cdiv(kv_len, NUM_KV_SPLITS)
    sp_start = kv_per_split * split_id
    sp_end = tl.minimum(sp_start + kv_per_split, kv_len)
    offs_c = tl.arange(0, BLOCK_C)
    offs_r = tl.arange(0, BLOCK_R)
    q_base = q_tok * STRIDE_QT
    q_nope = tl.load(Q + q_base + heads[:, None] * STRIDE_QH + offs_c[None, :]).to(FP8_DTYPE)
    q_rope = tl.load(Q + q_base + heads[:, None] * STRIDE_QH + (512 + offs_r[None, :])).to(FP8_DTYPE)
    e_max = tl.full([BLOCK_H], value=float("-inf"), dtype=tl.float32)
    e_sum = tl.zeros([BLOCK_H], dtype=tl.float32)
    acc = tl.zeros([BLOCK_H, BLOCK_C], dtype=tl.float32)
    num_tokens = sp_end - sp_start
    num_full = num_tokens // BLOCK_N
    full_end = sp_start + num_full * BLOCK_N

    for start_n in range(sp_start, full_end, BLOCK_N):
        kv_idx = kv_s + start_n + tl.arange(0, BLOCK_N)
        kv_nope = tl.load(KV + kv_idx[:, None] * STRIDE_KVT + offs_c[None, :],
                          cache_modifier=".cg")
        kv_rope = tl.load(KV + kv_idx[:, None] * STRIDE_KVT + (512 + offs_r[None, :]),
                          cache_modifier=".cg")
        k_nope = tl.trans(kv_nope)
        k_rope = tl.trans(kv_rope)
        qk = tl.dot(q_nope, k_nope) + tl.dot(q_rope, k_rope)
        qk = qk * score_scale
        new_max = tl.maximum(tl.max(qk, 1), e_max)
        rescale = tl.math.exp2((e_max - new_max) * LOG2E)
        p = tl.math.exp2((qk - new_max[:, None]) * LOG2E)
        acc = acc * rescale[:, None]
        acc = acc + tl.dot(p.to(FP8_DTYPE), kv_nope)
        e_sum = e_sum * rescale + tl.sum(p, 1)
        e_max = new_max

    if full_end < sp_end:
        offs_n = full_end + tl.arange(0, BLOCK_N)
        mask_n = offs_n < sp_end
        kv_idx = kv_s + offs_n
        kv_nope = tl.load(KV + kv_idx[:, None] * STRIDE_KVT + offs_c[None, :],
                          mask=mask_n[:, None], other=0.0, cache_modifier=".cg")
        kv_rope = tl.load(KV + kv_idx[:, None] * STRIDE_KVT + (512 + offs_r[None, :]),
                          mask=mask_n[:, None], other=0.0, cache_modifier=".cg")
        k_nope = tl.trans(kv_nope)
        k_rope = tl.trans(kv_rope)
        qk = tl.dot(q_nope, k_nope) + tl.dot(q_rope, k_rope)
        qk = qk * score_scale
        qk = tl.where(mask_n[None, :], qk, float("-inf"))
        new_max = tl.maximum(tl.max(qk, 1), e_max)
        rescale = tl.math.exp2((e_max - new_max) * LOG2E)
        p = tl.math.exp2((qk - new_max[:, None]) * LOG2E)
        acc = acc * rescale[:, None]
        acc = acc + tl.dot(p.to(FP8_DTYPE), kv_nope)
        e_sum = e_sum * rescale + tl.sum(p, 1)
        e_max = new_max

    safe_esum = tl.where(e_sum > 0, e_sum, 1.0)
    inv_esum = _fast_rcp(safe_esum)
    result = acc * inv_esum[:, None]
    if DIRECT_OUT:
        o_base = q_tok * STRIDE_OT
        tl.store(Out + o_base + heads[:, None] * STRIDE_OH + offs_c[None, :],
                 (result * kv_scale).to(tl.bfloat16))
    else:
        STRIDE_ACC: tl.constexpr = 512
        stride_acc_h = NUM_KV_SPLITS * STRIDE_ACC
        stride_acc_b = BLOCK_H * stride_acc_h
        acc_base = cur_batch * stride_acc_b + heads * stride_acc_h + split_id * STRIDE_ACC
        tl.store(Mid_Acc + acc_base[:, None] + offs_c[None, :], result)
        lse = tl.where(e_sum > 0, e_max + tl.math.log2(e_sum) * LN2, float("-inf"))
        stride_lse_h = NUM_KV_SPLITS
        stride_lse_b = BLOCK_H * stride_lse_h
        lse_base = cur_batch * stride_lse_b + heads * stride_lse_h + split_id
        tl.store(Mid_LSE + lse_base, lse)


@triton.jit
def _stage2_kernel(
    Mid_Acc, Mid_LSE, Out, qo_indptr, kv_scale_ptr,
    NUM_KV_SPLITS: tl.constexpr, BLOCK_DV: tl.constexpr,
    batch: tl.constexpr, GRID_SIZE: tl.constexpr,
    STRIDE_OT: tl.constexpr, STRIDE_OH: tl.constexpr,
    BLOCK_H: tl.constexpr,
):
    raw_pid = tl.program_id(0)
    pid = _remap_xcd(raw_pid, GRID_SIZE)
    cur_batch = pid % batch
    cur_head = pid // batch
    kv_scale = tl.load(kv_scale_ptr).to(tl.float32)
    q_tok = tl.load(qo_indptr + cur_batch)
    offs_d = tl.arange(0, BLOCK_DV)
    e_max = float("-inf")
    e_sum = 0.0
    acc = tl.zeros([BLOCK_DV], dtype=tl.float32)
    STRIDE_ACC: tl.constexpr = 512
    stride_acc_h = NUM_KV_SPLITS * STRIDE_ACC
    stride_acc_b = BLOCK_H * stride_acc_h
    acc_base = cur_batch * stride_acc_b + cur_head * stride_acc_h
    stride_lse_h = NUM_KV_SPLITS
    stride_lse_b = BLOCK_H * stride_lse_h
    lse_base = cur_batch * stride_lse_b + cur_head * stride_lse_h
    for s in range(NUM_KV_SPLITS):
        tv = tl.load(Mid_Acc + acc_base + s * STRIDE_ACC + offs_d)
        lse = tl.load(Mid_LSE + lse_base + s)
        new_max = tl.maximum(lse, e_max)
        old_scale = tl.math.exp2((e_max - new_max) * LOG2E)
        exp_lse = tl.math.exp2((lse - new_max) * LOG2E)
        acc = acc * old_scale + exp_lse * tv
        e_sum = e_sum * old_scale + exp_lse
        e_max = new_max
    safe_esum = tl.maximum(e_sum, 1e-12)
    inv_esum = _fast_rcp(safe_esum)
    result = acc * inv_esum * kv_scale
    o_base = q_tok * STRIDE_OT + cur_head * STRIDE_OH
    tl.store(Out + o_base + offs_d, result.to(tl.bfloat16))


def _get_splits(batch_size, kv_len):
    key = (batch_size, kv_len)
    if key in _SHAPE_CONFIGS:
        return _SHAPE_CONFIGS[key]
    max_tiles = max(1, kv_len // 64)
    if batch_size >= 256:
        return min(2, max_tiles)
    elif batch_size >= 64:
        return min(max(1, 512 // batch_size), max_tiles)
    else:
        splits = min(max(1, 768 // batch_size), max_tiles)
        while splits > 1 and batch_size * splits > 912:
            splits -= 1
        return splits


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    kv_fp8, kv_scale_t = kv_data["fp8"]
    kv = kv_fp8.view(torch.float8_e4m3fn).view(-1, QK_HEAD_DIM)
    batch_size = config["batch_size"]
    sm_scale = config["sm_scale"]
    total_q = q.shape[0]
    kv_len = kv.shape[0] // batch_size if batch_size > 0 else 0

    num_splits = _get_splits(batch_size, kv_len)

    dev = q.device
    direct_out = (num_splits == 1)

    if not direct_out:
        mid_acc = torch.empty(
            (batch_size, NUM_HEADS, num_splits, KV_LORA_RANK),
            dtype=torch.float32, device=dev,
        )
        mid_lse = torch.empty(
            (batch_size, NUM_HEADS, num_splits),
            dtype=torch.float32, device=dev,
        )
    else:
        mid_acc = torch.empty(1, dtype=torch.float32, device=dev)
        mid_lse = torch.empty(1, dtype=torch.float32, device=dev)

    out = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)

    grid_size1 = batch_size * num_splits
    _stage1_kernel[(grid_size1,)](
        q, kv, mid_acc, mid_lse, out, qo_indptr, kv_indptr,
        sm_scale, kv_scale_t,
        NUM_KV_SPLITS=num_splits, BLOCK_N=64, BLOCK_H=16,
        BLOCK_C=512, BLOCK_R=64, DIRECT_OUT=direct_out,
        GRID_SIZE=grid_size1,
        STRIDE_QT=NUM_HEADS * QK_HEAD_DIM,
        STRIDE_QH=QK_HEAD_DIM,
        STRIDE_KVT=QK_HEAD_DIM,
        STRIDE_OT=NUM_HEADS * V_HEAD_DIM,
        STRIDE_OH=V_HEAD_DIM,
        num_warps=4, num_stages=2, waves_per_eu=2,
    )

    if not direct_out:
        grid_size2 = NUM_HEADS * batch_size
        _stage2_kernel[(grid_size2,)](
            mid_acc, mid_lse, out, qo_indptr, kv_scale_t,
            NUM_KV_SPLITS=num_splits, BLOCK_DV=512,
            batch=batch_size, GRID_SIZE=grid_size2,
            STRIDE_OT=NUM_HEADS * V_HEAD_DIM,
            STRIDE_OH=V_HEAD_DIM,
            BLOCK_H=16,
            num_warps=4, num_stages=2,
        )

    return out
scrolls · 267 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 745692.

- """MLA decode Triton v135: constexpr strides.
+ """MLA decode Triton v140: v135 + separated LSE tensor (stride=512).
- All tensor strides are known constants. Making them constexprs:
- - Enables shift+add instead of multiply for address calculations
- (stride_kvt=576=2^9+2^6, used 128x per tile)
- - Eliminates 8 SGPR loads at kernel entry
- - Reduces kernel arguments from 16→8 (stage1), 9→4 (stage2)
- - Compiler can fully resolve Q/output address patterns at compile time
+ Isolates the separated tensor layout from v138 (which also changed
+ stage2 nw=2). Keeps stage2 nw=4 to avoid the bs32/kv8k regression.
- Based on v128 (GM=32.5us).
+ mid_acc: [batch, heads, splits, 512] stride=512 (power of 2, shift)
+ mid_lse: [batch, heads, splits] separate small tensor
"""
import torch
⋯ 44 unchanged lines
@triton.jit
def _stage1_kernel(
- Q, KV, Mid_O, Out,
+ Q, KV, Mid_Acc, Mid_LSE, Out,
qo_indptr, kv_indptr,
sm_scale, kv_scale_ptr,
NUM_KV_SPLITS: tl.constexpr, BLOCK_N: tl.constexpr,
⋯ 2 unchanged lines
STRIDE_QT: tl.constexpr, STRIDE_QH: tl.constexpr,
STRIDE_KVT: tl.constexpr,
STRIDE_OT: tl.constexpr, STRIDE_OH: tl.constexpr,
- STRIDE_MS: tl.constexpr,
):
raw_pid = tl.program_id(0)
pid = _remap_xcd(raw_pid, GRID_SIZE)
⋯ 68 unchanged lines
tl.store(Out + o_base + heads[:, None] * STRIDE_OH + offs_c[None, :],
(result * kv_scale).to(tl.bfloat16))
else:
- stride_mh = NUM_KV_SPLITS * STRIDE_MS
- stride_mb = BLOCK_H * stride_mh
- mid_base = cur_batch * stride_mb + heads * stride_mh + split_id * STRIDE_MS
- tl.store(Mid_O + mid_base[:, None] + offs_c[None, :], result)
+ STRIDE_ACC: tl.constexpr = 512
+ stride_acc_h = NUM_KV_SPLITS * STRIDE_ACC
+ stride_acc_b = BLOCK_H * stride_acc_h
+ acc_base = cur_batch * stride_acc_b + heads * stride_acc_h + split_id * STRIDE_ACC
+ tl.store(Mid_Acc + acc_base[:, None] + offs_c[None, :], result)
lse = tl.where(e_sum > 0, e_max + tl.math.log2(e_sum) * LN2, float("-inf"))
- tl.store(Mid_O + mid_base + 512, lse)
+ stride_lse_h = NUM_KV_SPLITS
+ stride_lse_b = BLOCK_H * stride_lse_h
+ lse_base = cur_batch * stride_lse_b + heads * stride_lse_h + split_id
+ tl.store(Mid_LSE + lse_base, lse)
@triton.jit
def _stage2_kernel(
- Mid_O, Out, qo_indptr, kv_scale_ptr,
+ Mid_Acc, Mid_LSE, Out, qo_indptr, kv_scale_ptr,
NUM_KV_SPLITS: tl.constexpr, BLOCK_DV: tl.constexpr,
batch: tl.constexpr, GRID_SIZE: tl.constexpr,
STRIDE_OT: tl.constexpr, STRIDE_OH: tl.constexpr,
- STRIDE_MS: tl.constexpr, BLOCK_H: tl.constexpr,
+ BLOCK_H: tl.constexpr,
):
raw_pid = tl.program_id(0)
pid = _remap_xcd(raw_pid, GRID_SIZE)
⋯ 5 unchanged lines
e_max = float("-inf")
e_sum = 0.0
acc = tl.zeros([BLOCK_DV], dtype=tl.float32)
- stride_mh = NUM_KV_SPLITS * STRIDE_MS
- stride_mb = BLOCK_H * stride_mh
- mid_base = cur_batch * stride_mb + cur_head * stride_mh
+ STRIDE_ACC: tl.constexpr = 512
+ stride_acc_h = NUM_KV_SPLITS * STRIDE_ACC
+ stride_acc_b = BLOCK_H * stride_acc_h
+ acc_base = cur_batch * stride_acc_b + cur_head * stride_acc_h
+ stride_lse_h = NUM_KV_SPLITS
+ stride_lse_b = BLOCK_H * stride_lse_h
+ lse_base = cur_batch * stride_lse_b + cur_head * stride_lse_h
for s in range(NUM_KV_SPLITS):
- tv = tl.load(Mid_O + mid_base + s * STRIDE_MS + offs_d)
- lse = tl.load(Mid_O + mid_base + s * STRIDE_MS + 512)
+ tv = tl.load(Mid_Acc + acc_base + s * STRIDE_ACC + offs_d)
+ lse = tl.load(Mid_LSE + lse_base + s)
new_max = tl.maximum(lse, e_max)
old_scale = tl.math.exp2((e_max - new_max) * LOG2E)
exp_lse = tl.math.exp2((lse - new_max) * LOG2E)
⋯ 36 unchanged lines
dev = q.device
direct_out = (num_splits == 1)
- mid_o = torch.empty(
- (batch_size, NUM_HEADS, num_splits, KV_LORA_RANK + 1),
- dtype=torch.float32, device=dev,
- ) if not direct_out else torch.empty(1, dtype=torch.float32, device=dev)
+ if not direct_out:
+ mid_acc = torch.empty(
+ (batch_size, NUM_HEADS, num_splits, KV_LORA_RANK),
+ dtype=torch.float32, device=dev,
+ )
+ mid_lse = torch.empty(
+ (batch_size, NUM_HEADS, num_splits),
+ dtype=torch.float32, device=dev,
+ )
+ else:
+ mid_acc = torch.empty(1, dtype=torch.float32, device=dev)
+ mid_lse = torch.empty(1, dtype=torch.float32, device=dev)
+
out = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)
grid_size1 = batch_size * num_splits
_stage1_kernel[(grid_size1,)](
- q, kv, mid_o, out, qo_indptr, kv_indptr,
+ q, kv, mid_acc, mid_lse, out, qo_indptr, kv_indptr,
sm_scale, kv_scale_t,
NUM_KV_SPLITS=num_splits, BLOCK_N=64, BLOCK_H=16,
BLOCK_C=512, BLOCK_R=64, DIRECT_OUT=direct_out,
⋯ 3 unchanged lines
STRIDE_KVT=QK_HEAD_DIM,
STRIDE_OT=NUM_HEADS * V_HEAD_DIM,
STRIDE_OH=V_HEAD_DIM,
- STRIDE_MS=KV_LORA_RANK + 1,
num_warps=4, num_stages=2, waves_per_eu=2,
)
if not direct_out:
grid_size2 = NUM_HEADS * batch_size
_stage2_kernel[(grid_size2,)](
- mid_o, out, qo_indptr, kv_scale_t,
+ mid_acc, mid_lse, out, qo_indptr, kv_scale_t,
NUM_KV_SPLITS=num_splits, BLOCK_DV=512,
batch=batch_size, GRID_SIZE=grid_size2,
STRIDE_OT=NUM_HEADS * V_HEAD_DIM,
STRIDE_OH=V_HEAD_DIM,
- STRIDE_MS=KV_LORA_RANK + 1,
BLOCK_H=16,
num_warps=4, num_stages=2,
)
scrolls · 142 diff lines total

Best evidence level for this revision: reported

JSON