Skip to content
KernelIndex
Search⌘K

submission 755102

Law1912 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

svm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-755102?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
#22 of 766
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:7ce405b628a401a480ca83f433a839b28572c1f9d5beafc3ef7287ef6d7fb52d
license declaredunknown
license concludedunknown
authorsLaw1912
imported2026-08-15

Techniques

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

fp8FP8 = tl.float8e4nv
mmalogits = tl.dot(q_c, tl.trans(kc)) + tl.dot(q_r, tl.trans(kr))
num-warps = 4num_warps=4, num_stages=2, waves_per_eu=2,
stages = 2num_warps=4, num_stages=2, waves_per_eu=2,
tile-k = 64NH=_H, DC=_DL, DR=64, BK=64,

Kernel source

svm.py238 lines
"""Two-pass grouped FP8 attention for decode with log-sum-exp merging."""

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

_H, _DQ, _DL, _DV = 16, 576, 512, 512

_PARTITION_MAP = [
    (4, 1024, 16), (4, 8192, 32),
    (32, 1024, 4), (32, 8192, 8),
    (64, 1024, 4), (64, 8192, 8),
    (256, 1024, 1), (256, 8192, 2),
]


def _select_partitions(batch, seqlen):
    for b, s, p in _PARTITION_MAP:
        if b == batch and s == seqlen:
            return p
    limit = max(1, seqlen // 64)
    if batch >= 256:
        return min(2, limit)
    if batch >= 64:
        return min(max(1, 512 // batch), limit)
    p = min(max(1, 768 // batch), limit)
    while p > 1 and batch * p > 912:
        p -= 1
    return p


@triton.jit
def _compute_chunks(
    q_ptr, kv_ptr, acc_buf, lse_buf, final_buf,
    token_map, kv_spans,
    scale_f, scale_kv_ptr,
    qs0: tl.constexpr, qs1: tl.constexpr,
    kvs0: tl.constexpr,
    os0: tl.constexpr, os1: tl.constexpr,
    NH: tl.constexpr, DC: tl.constexpr, DR: tl.constexpr, BK: tl.constexpr,
    NP: tl.constexpr, NITERS: tl.constexpr, FUSE_OUT: tl.constexpr,
    NWGS: tl.constexpr,
    NUM_XCD: tl.constexpr = 8,
):
    raw = tl.program_id(0)
    wgs_per = (NWGS + NUM_XCD - 1) // NUM_XCD
    n_full = NWGS % NUM_XCD
    n_full = NUM_XCD if n_full == 0 else n_full
    xid = raw % NUM_XCD
    lid = raw // NUM_XCD
    if xid < n_full:
        gid = xid * wgs_per + lid
    else:
        gid = n_full * wgs_per + (xid - n_full) * (wgs_per - 1) + lid

    sample = gid // NP
    part = gid % NP

    ROPE_BASE: tl.constexpr = 512
    L2E: tl.constexpr = 1.4426950408889634
    LN2_VAL: tl.constexpr = 0.6931471805599453
    FP8 = tl.float8e4nv

    ks = tl.load(scale_kv_ptr).to(tl.float32)
    combined = scale_f * ks * L2E

    tok = tl.load(token_map + sample)
    kv0 = tl.load(kv_spans + sample)
    kv1 = tl.load(kv_spans + sample + 1)
    total = kv1 - kv0
    per_p = tl.cdiv(total, NP)
    origin = per_p * part

    hh = tl.arange(0, NH)
    cc = tl.arange(0, DC)
    rr = tl.arange(0, DR)

    qoff = tok * qs0
    q_c = tl.load(q_ptr + qoff + hh[:, None] * qs1 + cc[None, :]).to(FP8)
    q_r = tl.load(q_ptr + qoff + hh[:, None] * qs1 + (ROPE_BASE + rr[None, :])).to(FP8)

    mx = tl.full([NH], value=float("-inf"), dtype=tl.float32)
    sm = tl.zeros([NH], dtype=tl.float32)
    ov = tl.zeros([NH, DC], dtype=tl.float32)

    loop_n = NITERS if NITERS > 0 else (tl.minimum(origin + per_p, total) - origin) // BK
    for i in range(loop_n):
        idx = kv0 + origin + i * BK + tl.arange(0, BK)
        kc = tl.load(kv_ptr + idx[:, None] * kvs0 + cc[None, :], cache_modifier=".cg")
        kr = tl.load(kv_ptr + idx[:, None] * kvs0 + (ROPE_BASE + rr[None, :]), cache_modifier=".cg")

        logits = tl.dot(q_c, tl.trans(kc)) + tl.dot(q_r, tl.trans(kr))
        logits *= combined

        new_mx = tl.maximum(tl.max(logits, 1), mx)
        alpha = tl.math.exp2(mx - new_mx)
        beta = tl.math.exp2(logits - new_mx[:, None])

        ov = ov * alpha[:, None] + tl.dot(beta.to(FP8), kc)
        sm = sm * alpha + tl.sum(beta, 1)
        mx = new_mx

    guard = tl.where(sm > 0, sm, 1.0)
    rcp = tl.inline_asm_elementwise("v_rcp_f32_e32 $0, $1", "=v, v",
                                    [guard], dtype=tl.float32,
                                    is_pure=True, pack=1)
    normalized = ov * rcp[:, None]

    if FUSE_OUT:
        dst = tok * os0
        tl.store(final_buf + dst + hh[:, None] * os1 + cc[None, :],
                 (normalized * ks).to(tl.bfloat16))
    else:
        BW: tl.constexpr = 512
        s_part = BW
        s_head = NP * s_part
        s_batch = NH * s_head
        a = sample * s_batch + hh * s_head + part * s_part
        tl.store(acc_buf + a[:, None] + cc[None, :], normalized)

        lse_val = tl.where(sm > 0,
                           (mx + tl.math.log2(sm)) * LN2_VAL,
                           float("-inf"))
        la = sample * (NH * NP) + hh * NP + part
        tl.store(lse_buf + la, lse_val)


@triton.jit
def _merge_chunks(
    acc_buf, lse_buf, final_buf, token_map, scale_kv_ptr,
    NP: tl.constexpr, DV: tl.constexpr,
    NB: tl.constexpr, NWGS: tl.constexpr,
    os0: tl.constexpr, os1: tl.constexpr,
    NH: tl.constexpr,
    NUM_XCD: tl.constexpr = 8,
):
    L2E: tl.constexpr = 1.4426950408889634
    raw = tl.program_id(0)
    wgs_per = (NWGS + NUM_XCD - 1) // NUM_XCD
    n_full = NWGS % NUM_XCD
    n_full = NUM_XCD if n_full == 0 else n_full
    xid = raw % NUM_XCD
    lid = raw // NUM_XCD
    if xid < n_full:
        gid = xid * wgs_per + lid
    else:
        gid = n_full * wgs_per + (xid - n_full) * (wgs_per - 1) + lid

    b = gid % NB
    h = gid // NB
    ks = tl.load(scale_kv_ptr).to(tl.float32)
    tok = tl.load(token_map + b)
    dd = tl.arange(0, DV)

    BW: tl.constexpr = 512
    s_part = BW
    s_head = NP * s_part
    s_batch = NH * s_head
    data_off = b * s_batch + h * s_head
    lse_off = b * (NH * NP) + h * NP

    best = float("-inf")
    wsum = 0.0
    result = tl.zeros([DV], dtype=tl.float32)

    for k in range(NP):
        chunk_v = tl.load(acc_buf + data_off + k * s_part + dd)
        chunk_l = tl.load(lse_buf + lse_off + k)
        nb = tl.maximum(chunk_l, best)
        wa = tl.math.exp2((best - nb) * L2E)
        wb = tl.math.exp2((chunk_l - nb) * L2E)
        result = result * wa + wb * chunk_v
        wsum = wsum * wa + wb
        best = nb

    clamped = tl.maximum(wsum, 1e-12)
    inv = tl.inline_asm_elementwise("v_rcp_f32_e32 $0, $1", "=v, v",
                                    [clamped], dtype=tl.float32,
                                    is_pure=True, pack=1)
    base = tok * os0 + h * os1
    tl.store(final_buf + base + dd, (result * inv * ks).to(tl.bfloat16))


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    kv_raw, kv_sc = kv_data["fp8"]
    kv = kv_raw.view(torch.float8_e4m3fn).view(-1, _DQ)

    B = config["batch_size"]
    sm = config["sm_scale"]
    T = q.shape[0]
    S = kv.shape[0] // B if B > 0 else 0

    P = _select_partitions(B, S)
    fused = P == 1

    chunk_len = (S + P - 1) // P
    fixed = chunk_len // 64
    if chunk_len % 64 != 0 or fixed > 4:
        fixed = 0

    output = torch.empty((T, _H, _DV), dtype=torch.bfloat16, device=q.device)

    if fused:
        ab = torch.empty(1, dtype=torch.float32, device=q.device)
        lb = torch.empty(1, dtype=torch.float32, device=q.device)
    else:
        ab = torch.empty((B, _H, P, _DL), dtype=torch.float32, device=q.device)
        lb = torch.empty((B, _H, P), dtype=torch.float32, device=q.device)

    g1 = B * P
    _compute_chunks[(g1,)](
        q, kv, ab, lb, output,
        qo_indptr, kv_indptr,
        sm, kv_sc,
        qs0=_H * _DQ, qs1=_DQ,
        kvs0=_DQ,
        os0=_H * _DV, os1=_DV,
        NH=_H, DC=_DL, DR=64, BK=64,
        NP=P, NITERS=fixed, FUSE_OUT=fused,
        NWGS=g1,
        num_warps=4, num_stages=2, waves_per_eu=2,
    )

    if not fused:
        g2 = _H * B
        _merge_chunks[(g2,)](
            ab, lb, output, qo_indptr, kv_sc,
            NP=P, DV=_DL,
            NB=B, NWGS=g2,
            os0=_H * _DV, os1=_DV,
            NH=_H,
            num_warps=4, num_stages=2,
        )

    return output
scrolls · 238 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