Skip to content
KernelIndex
Search⌘K

submission 589795

Ananda Sai A · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:7d9d1731d316d9031826446585a696594d14fdeb3b90ebcdc9718252d0971a99
license declaredunknown
license concludedunknown
authorsAnanda Sai A
imported2026-08-15

Techniques

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

mmaqk = tl.dot(qn, kn16) + tl.dot(qr, kr16)
num-warps = 4num_warps=4, num_stages=1, **ex,
stages = 1num_warps=4, num_stages=1, **ex,

Kernel source

submission_v15_ultimate.py289 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
Ultimate MLA decode: best path per shape.
- bmm bf16: small shapes (bs<=4 all, bs<=64 kv<=1024)
- Triton fp8 flash-decode: bs=32/kv=8192 (beats AITER by 30%)
- AITER fp8: large shapes (bs>=64/kv=8192, bs=256/kv=1024)
"""
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
import math
from task import input_t, output_t

from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

NUM_Q_HEADS = 16
NUM_KV_HEADS = 1
QK_DIM = 576
V_DIM = 512
SM_SCALE = 1.0 / math.sqrt(576)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8

_meta_cache = {}
_idx_cache = {}
_out_cache = {}
_buf = {}
_warm = set()


# ═══════════════════════════════════════════════
# Triton flash-decode fp8 (for medium shapes)
# ═══════════════════════════════════════════════

@triton.jit
def _flash_s1(
    Q, KV_FP8, kv_scale_ptr, sm_scale,
    kv_indptr, Att_Out, Att_Lse,
    stride_qb, stride_qh, stride_kv_tok,
    stride_ab, stride_ah, stride_as,
    stride_lb, stride_lh,
    BLOCK_N: tl.constexpr, BLOCK_H: tl.constexpr,
    NUM_SPLITS: tl.constexpr,
    BLOCK_NOPE: tl.constexpr, BLOCK_ROPE: tl.constexpr,
    BLOCK_DV: tl.constexpr,
    Lnope: tl.constexpr, Lrope: tl.constexpr, Lv: tl.constexpr,
):
    bid = tl.program_id(0)
    sid = tl.program_id(1)

    heads = tl.arange(0, BLOCK_H)
    mask_h = heads < 16
    o_nope = tl.arange(0, BLOCK_NOPE)
    o_rope = tl.arange(0, BLOCK_ROPE)
    o_rope_s = Lnope + o_rope
    o_dv = tl.arange(0, BLOCK_DV)
    mn = o_nope < Lnope
    mr = o_rope < Lrope
    mv = o_dv < Lv

    ks = tl.load(kv_indptr + bid)
    ke = tl.load(kv_indptr + bid + 1)
    kl = ke - ks

    ss = tl.cdiv(kl, NUM_SPLITS)
    ss = tl.cdiv(ss, BLOCK_N) * BLOCK_N
    ms = sid * ss
    me = tl.minimum(ms + ss, kl)

    emax = tl.zeros([BLOCK_H], dtype=tl.float32) - float("inf")
    esum = tl.zeros([BLOCK_H], dtype=tl.float32)
    acc = tl.zeros([BLOCK_H, BLOCK_DV], dtype=tl.float32)
    sc = tl.load(kv_scale_ptr)

    if me > ms:
        qb = bid * stride_qb
        qn = tl.load(Q + qb + heads[:, None] * stride_qh + o_nope[None, :],
                      mask=mask_h[:, None] & mn[None, :], other=0.0).to(tl.float16)
        qr = tl.load(Q + qb + heads[:, None] * stride_qh + o_rope_s[None, :],
                      mask=mask_h[:, None] & mr[None, :], other=0.0).to(tl.float16)

        for t in range(ms, me, BLOCK_N):
            no = tl.arange(0, BLOCK_N)
            nm = (t + no) < me
            ti = (ks + t + no) * stride_kv_tok

            kn = tl.load(KV_FP8 + ti[None, :] + o_nope[:, None],
                         mask=nm[None, :] & mn[:, None], other=0.0)
            kn16 = (kn.to(tl.float32) * sc).to(tl.float16)

            kr = tl.load(KV_FP8 + ti[None, :] + o_rope_s[:, None],
                         mask=nm[None, :] & mr[:, None], other=0.0)
            kr16 = (kr.to(tl.float32) * sc).to(tl.float16)

            qk = tl.dot(qn, kn16) + tl.dot(qr, kr16)
            qk = qk.to(tl.float32) * sm_scale
            qk = tl.where(mask_h[:, None] & nm[None, :], qk, float("-inf"))

            vf = tl.load(KV_FP8 + ti[:, None] + o_dv[None, :],
                         mask=nm[:, None] & mv[None, :], other=0.0)
            v16 = (vf.to(tl.float32) * sc).to(tl.float16)

            ne = tl.maximum(tl.max(qk, 1), emax)
            rs = tl.exp(emax - ne)
            p = tl.exp(qk - ne[:, None])
            acc = acc * rs[:, None] + tl.dot(p.to(tl.float16), v16).to(tl.float32)
            esum = esum * rs + tl.sum(p, 1)
            emax = ne

    ob = bid * stride_ab + heads[:, None] * stride_ah + sid * stride_as + o_dv[None, :]
    tl.store(Att_Out + ob, acc / tl.maximum(esum[:, None], 1e-12), mask=mask_h[:, None] & mv[None, :])
    lb = bid * stride_lb + heads * stride_lh + sid
    tl.store(Att_Lse + lb, emax + tl.log(tl.maximum(esum, 1e-12)), mask=mask_h)


@triton.jit
def _flash_s2(
    Att_Out, Att_Lse, O,
    stride_ab, stride_ah, stride_as,
    stride_lb, stride_lh,
    stride_ob, stride_oh,
    NS: tl.constexpr, BDV: tl.constexpr, Lv: tl.constexpr,
):
    bid = tl.program_id(0)
    hid = tl.program_id(1)
    od = tl.arange(0, BDV)
    md = od < Lv
    em = -float("inf")
    es = 0.0
    ac = tl.zeros([BDV], dtype=tl.float32)
    for s in range(NS):
        l = tl.load(Att_Lse + bid * stride_lb + hid * stride_lh + s)
        if l > -1e30:
            pv = tl.load(Att_Out + bid * stride_ab + hid * stride_ah + s * stride_as + od, mask=md, other=0.0)
            nm = tl.maximum(l, em)
            os = tl.exp(em - nm)
            ns = tl.exp(l - nm)
            ac = ac * os + ns * pv
            es = es * os + ns
            em = nm
    tl.store(O + bid * stride_ob + hid * stride_oh + od, (ac / tl.maximum(es, 1e-12)).to(tl.bfloat16), mask=md)


def _triton_fp8(q, kv_data, kv_indptr, config, nsplits):
    bs = config["batch_size"]
    kv_fp8, kv_scale = kv_data["fp8"]
    kv_flat = kv_fp8.view(-1, QK_DIM)
    q_r = q.view(bs, NUM_Q_HEADS, QK_DIM)
    k = (bs, nsplits)
    if k not in _buf:
        d = q.device
        _buf[k] = (
            torch.empty((bs, NUM_Q_HEADS, nsplits, V_DIM), dtype=torch.float32, device=d),
            torch.empty((bs, NUM_Q_HEADS, nsplits), dtype=torch.float32, device=d),
            torch.empty((bs, NUM_Q_HEADS, V_DIM), dtype=torch.bfloat16, device=d),
        )
    ao, al, o = _buf[k]
    BN = 32
    ex = {}
    try:
        if triton.runtime.driver.active.get_current_target().backend == "hip":
            ex = {"waves_per_eu": 1, "matrix_instr_nonkdim": 16, "kpack": 2}
    except Exception:
        pass
    _flash_s1[(bs, nsplits)](
        q_r, kv_flat, kv_scale, SM_SCALE, kv_indptr, ao, al,
        q_r.stride(0), q_r.stride(1), kv_flat.stride(0),
        ao.stride(0), ao.stride(1), ao.stride(2),
        al.stride(0), al.stride(1),
        BLOCK_N=BN, BLOCK_H=16, NUM_SPLITS=nsplits,
        BLOCK_NOPE=512, BLOCK_ROPE=64, BLOCK_DV=512,
        Lnope=512, Lrope=64, Lv=512,
        num_warps=4, num_stages=1, **ex,
    )
    _flash_s2[(bs, NUM_Q_HEADS)](
        ao, al, o,
        ao.stride(0), ao.stride(1), ao.stride(2),
        al.stride(0), al.stride(1),
        o.stride(0), o.stride(1),
        NS=nsplits, BDV=512, Lv=512,
        num_warps=4, num_stages=1, **ex,
    )
    return o


# ═══════════════════════════════════════════════
# AITER fp8 (for largest shapes)
# ═══════════════════════════════════════════════

def _quantize_fp8(tensor):
    finfo = torch.finfo(FP8_DTYPE)
    amax = tensor.abs().amax().clamp(min=1e-12)
    scale = amax / finfo.max
    return (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE), scale.float().reshape(1)

def _aiter_fp8(q, kv_data, qo_indptr, kv_indptr, config, num_splits):
    bs = config["batch_size"]
    q_len = config["q_seq_len"]
    total_kv = int(kv_indptr[-1].item())
    total_q = q.shape[0]

    q_fp8, q_scale = _quantize_fp8(q)
    kv_fp8, kv_scale = kv_data["fp8"]
    kv_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, NUM_KV_HEADS, kv_fp8.shape[-1])

    key = (bs, num_splits, str(q_fp8.dtype), str(kv_fp8.dtype))
    if key not in _meta_cache:
        info = get_mla_metadata_info_v1(bs, q_len, NUM_Q_HEADS, q_fp8.dtype, kv_fp8.dtype,
            is_sparse=False, fast_mode=False, num_kv_splits=num_splits, intra_batch_mode=True)
        _meta_cache[key] = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
    work = _meta_cache[key]
    kv_last = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
    get_mla_metadata_v1(qo_indptr, kv_indptr, kv_last,
        NUM_Q_HEADS, NUM_KV_HEADS, True,
        work[0], work[2], work[1], work[3], work[4], work[5],
        page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16),
        max_seqlen_qo=q_len, uni_seqlen_qo=q_len,
        fast_mode=False, max_split_per_batch=num_splits,
        intra_batch_mode=True, dtype_q=q_fp8.dtype, dtype_kv=kv_fp8.dtype)

    if total_kv not in _idx_cache:
        _idx_cache[total_kv] = torch.arange(total_kv, dtype=torch.int32, device="cuda")
    if total_q not in _out_cache:
        _out_cache[total_q] = torch.empty((total_q, NUM_Q_HEADS, V_DIM), dtype=torch.bfloat16, device="cuda")

    o = _out_cache[total_q]
    mla_decode_fwd(q_fp8.view(-1, NUM_Q_HEADS, QK_DIM), kv_4d, o,
        qo_indptr, kv_indptr, _idx_cache[total_kv], kv_last, q_len,
        page_size=PAGE_SIZE, nhead_kv=NUM_KV_HEADS, sm_scale=SM_SCALE, logit_cap=0.0,
        num_kv_splits=num_splits, q_scale=q_scale, kv_scale=kv_scale,
        intra_batch_mode=True,
        work_meta_data=work[0], work_indptr=work[1], work_info_set=work[2],
        reduce_indptr=work[3], reduce_final_map=work[4], reduce_partial_map=work[5])
    return o


# ═══════════════════════════════════════════════
# BMM bf16 (for small shapes)
# ═══════════════════════════════════════════════

def _bmm(q, kv_data, config):
    bs = config["batch_size"]
    kv_len = config["kv_seq_len"]
    kv = kv_data["bf16"].view(bs, kv_len, QK_DIM)
    Q = q.view(bs, NUM_Q_HEADS, QK_DIM)
    V = kv[:, :, :V_DIM]
    s = torch.bmm(Q, kv.transpose(1, 2)) * SM_SCALE
    w = F.softmax(s, dim=-1, dtype=torch.float32).to(torch.bfloat16)
    return torch.bmm(w, V)


# ═══════════════════════════════════════════════
# Shape dispatch based on measured data
# ═══════════════════════════════════════════════

# Per-shape best path (from benchmarks):
# bmm: bs=4/1k(26), bs=4/8k(43), bs=32/1k(42), bs=64/1k(61)
# Triton fp8: bs=32/8k(120) [AITER was 166]
# AITER fp8: bs=64/8k(206), bs=256/1k(161), bs=256/8k(353)

_AITER_SPLITS = {
    (64, 8192): 32,
    (256, 1024): 16,
    (256, 8192): 32,
}

def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    bs = config["batch_size"]
    kv_len = config["kv_seq_len"]
    sk = (bs, kv_len)
    if sk not in _warm:
        _warm.add(sk)

    # bmm: small shapes
    if bs <= 4 or (bs <= 64 and kv_len <= 1024):
        return _bmm(q, kv_data, config)

    # Triton fp8: bs=32/kv=8192 (beats AITER)
    if bs <= 32 and kv_len <= 8192:
        return _triton_fp8(q, kv_data, kv_indptr, config, 8)

    # AITER fp8: large shapes
    return _aiter_fp8(q, kv_data, qo_indptr, kv_indptr, config, _AITER_SPLITS.get(sk, 32))
scrolls · 289 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 589553.

#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
+ """
+ Ultimate MLA decode: best path per shape.
+ - bmm bf16: small shapes (bs<=4 all, bs<=64 kv<=1024)
+ - Triton fp8 flash-decode: bs=32/kv=8192 (beats AITER by 30%)
+ - AITER fp8: large shapes (bs>=64/kv=8192, bs=256/kv=1024)
+ """
import torch
import torch.nn.functional as F
+ import triton
+ import triton.language as tl
+ import math
from task import input_t, output_t
from aiter.mla import mla_decode_fwd
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
- NUM_HEADS = 16
+ NUM_Q_HEADS = 16
NUM_KV_HEADS = 1
- QK_HEAD_DIM = 576
- V_HEAD_DIM = 512
- SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
+ QK_DIM = 576
+ V_DIM = 512
+ SM_SCALE = 1.0 / math.sqrt(576)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
+ _meta_cache = {}
+ _idx_cache = {}
+ _out_cache = {}
+ _buf = {}
_warm = set()
+
+ # ═══════════════════════════════════════════════
+ # Triton flash-decode fp8 (for medium shapes)
+ # ═══════════════════════════════════════════════
+
+ @triton.jit
+ def _flash_s1(
+ Q, KV_FP8, kv_scale_ptr, sm_scale,
+ kv_indptr, Att_Out, Att_Lse,
+ stride_qb, stride_qh, stride_kv_tok,
+ stride_ab, stride_ah, stride_as,
+ stride_lb, stride_lh,
+ BLOCK_N: tl.constexpr, BLOCK_H: tl.constexpr,
+ NUM_SPLITS: tl.constexpr,
+ BLOCK_NOPE: tl.constexpr, BLOCK_ROPE: tl.constexpr,
+ BLOCK_DV: tl.constexpr,
+ Lnope: tl.constexpr, Lrope: tl.constexpr, Lv: tl.constexpr,
+ ):
+ bid = tl.program_id(0)
+ sid = tl.program_id(1)
+
+ heads = tl.arange(0, BLOCK_H)
+ mask_h = heads < 16
+ o_nope = tl.arange(0, BLOCK_NOPE)
+ o_rope = tl.arange(0, BLOCK_ROPE)
+ o_rope_s = Lnope + o_rope
+ o_dv = tl.arange(0, BLOCK_DV)
+ mn = o_nope < Lnope
+ mr = o_rope < Lrope
+ mv = o_dv < Lv
+
+ ks = tl.load(kv_indptr + bid)
+ ke = tl.load(kv_indptr + bid + 1)
+ kl = ke - ks
+
+ ss = tl.cdiv(kl, NUM_SPLITS)
+ ss = tl.cdiv(ss, BLOCK_N) * BLOCK_N
+ ms = sid * ss
+ me = tl.minimum(ms + ss, kl)
+
+ emax = tl.zeros([BLOCK_H], dtype=tl.float32) - float("inf")
+ esum = tl.zeros([BLOCK_H], dtype=tl.float32)
+ acc = tl.zeros([BLOCK_H, BLOCK_DV], dtype=tl.float32)
+ sc = tl.load(kv_scale_ptr)
+
+ if me > ms:
+ qb = bid * stride_qb
+ qn = tl.load(Q + qb + heads[:, None] * stride_qh + o_nope[None, :],
+ mask=mask_h[:, None] & mn[None, :], other=0.0).to(tl.float16)
+ qr = tl.load(Q + qb + heads[:, None] * stride_qh + o_rope_s[None, :],
+ mask=mask_h[:, None] & mr[None, :], other=0.0).to(tl.float16)
+
+ for t in range(ms, me, BLOCK_N):
+ no = tl.arange(0, BLOCK_N)
+ nm = (t + no) < me
+ ti = (ks + t + no) * stride_kv_tok
+
+ kn = tl.load(KV_FP8 + ti[None, :] + o_nope[:, None],
+ mask=nm[None, :] & mn[:, None], other=0.0)
+ kn16 = (kn.to(tl.float32) * sc).to(tl.float16)
+
+ kr = tl.load(KV_FP8 + ti[None, :] + o_rope_s[:, None],
+ mask=nm[None, :] & mr[:, None], other=0.0)
+ kr16 = (kr.to(tl.float32) * sc).to(tl.float16)
+
+ qk = tl.dot(qn, kn16) + tl.dot(qr, kr16)
+ qk = qk.to(tl.float32) * sm_scale
+ qk = tl.where(mask_h[:, None] & nm[None, :], qk, float("-inf"))
+
+ vf = tl.load(KV_FP8 + ti[:, None] + o_dv[None, :],
+ mask=nm[:, None] & mv[None, :], other=0.0)
+ v16 = (vf.to(tl.float32) * sc).to(tl.float16)
+
+ ne = tl.maximum(tl.max(qk, 1), emax)
+ rs = tl.exp(emax - ne)
+ p = tl.exp(qk - ne[:, None])
+ acc = acc * rs[:, None] + tl.dot(p.to(tl.float16), v16).to(tl.float32)
+ esum = esum * rs + tl.sum(p, 1)
+ emax = ne
+
+ ob = bid * stride_ab + heads[:, None] * stride_ah + sid * stride_as + o_dv[None, :]
+ tl.store(Att_Out + ob, acc / tl.maximum(esum[:, None], 1e-12), mask=mask_h[:, None] & mv[None, :])
+ lb = bid * stride_lb + heads * stride_lh + sid
+ tl.store(Att_Lse + lb, emax + tl.log(tl.maximum(esum, 1e-12)), mask=mask_h)
+
+
+ @triton.jit
+ def _flash_s2(
+ Att_Out, Att_Lse, O,
+ stride_ab, stride_ah, stride_as,
+ stride_lb, stride_lh,
+ stride_ob, stride_oh,
+ NS: tl.constexpr, BDV: tl.constexpr, Lv: tl.constexpr,
+ ):
+ bid = tl.program_id(0)
+ hid = tl.program_id(1)
+ od = tl.arange(0, BDV)
+ md = od < Lv
+ em = -float("inf")
+ es = 0.0
+ ac = tl.zeros([BDV], dtype=tl.float32)
+ for s in range(NS):
+ l = tl.load(Att_Lse + bid * stride_lb + hid * stride_lh + s)
+ if l > -1e30:
+ pv = tl.load(Att_Out + bid * stride_ab + hid * stride_ah + s * stride_as + od, mask=md, other=0.0)
+ nm = tl.maximum(l, em)
+ os = tl.exp(em - nm)
+ ns = tl.exp(l - nm)
+ ac = ac * os + ns * pv
+ es = es * os + ns
+ em = nm
+ tl.store(O + bid * stride_ob + hid * stride_oh + od, (ac / tl.maximum(es, 1e-12)).to(tl.bfloat16), mask=md)
+
+
+ def _triton_fp8(q, kv_data, kv_indptr, config, nsplits):
+ bs = config["batch_size"]
+ kv_fp8, kv_scale = kv_data["fp8"]
+ kv_flat = kv_fp8.view(-1, QK_DIM)
+ q_r = q.view(bs, NUM_Q_HEADS, QK_DIM)
+ k = (bs, nsplits)
+ if k not in _buf:
+ d = q.device
+ _buf[k] = (
+ torch.empty((bs, NUM_Q_HEADS, nsplits, V_DIM), dtype=torch.float32, device=d),
+ torch.empty((bs, NUM_Q_HEADS, nsplits), dtype=torch.float32, device=d),
+ torch.empty((bs, NUM_Q_HEADS, V_DIM), dtype=torch.bfloat16, device=d),
+ )
+ ao, al, o = _buf[k]
+ BN = 32
+ ex = {}
+ try:
+ if triton.runtime.driver.active.get_current_target().backend == "hip":
+ ex = {"waves_per_eu": 1, "matrix_instr_nonkdim": 16, "kpack": 2}
+ except Exception:
+ pass
+ _flash_s1[(bs, nsplits)](
+ q_r, kv_flat, kv_scale, SM_SCALE, kv_indptr, ao, al,
+ q_r.stride(0), q_r.stride(1), kv_flat.stride(0),
+ ao.stride(0), ao.stride(1), ao.stride(2),
+ al.stride(0), al.stride(1),
+ BLOCK_N=BN, BLOCK_H=16, NUM_SPLITS=nsplits,
+ BLOCK_NOPE=512, BLOCK_ROPE=64, BLOCK_DV=512,
+ Lnope=512, Lrope=64, Lv=512,
+ num_warps=4, num_stages=1, **ex,
+ )
+ _flash_s2[(bs, NUM_Q_HEADS)](
+ ao, al, o,
+ ao.stride(0), ao.stride(1), ao.stride(2),
+ al.stride(0), al.stride(1),
+ o.stride(0), o.stride(1),
+ NS=nsplits, BDV=512, Lv=512,
+ num_warps=4, num_stages=1, **ex,
+ )
+ return o
+
+
+ # ═══════════════════════════════════════════════
+ # AITER fp8 (for largest shapes)
+ # ═══════════════════════════════════════════════
+
def _quantize_fp8(tensor):
finfo = torch.finfo(FP8_DTYPE)
amax = tensor.abs().amax().clamp(min=1e-12)
scale = amax / finfo.max
return (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE), scale.float().reshape(1)
- def _aiter_fp8_decode(q, kv_data, qo_indptr, kv_indptr, config, num_splits=32):
- """Standard AITER fp8 path for large shapes."""
+ def _aiter_fp8(q, kv_data, qo_indptr, kv_indptr, config, num_splits):
bs = config["batch_size"]
q_len = config["q_seq_len"]
total_kv = int(kv_indptr[-1].item())
+ total_q = q.shape[0]
q_fp8, q_scale = _quantize_fp8(q)
kv_fp8, kv_scale = kv_data["fp8"]
kv_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, NUM_KV_HEADS, kv_fp8.shape[-1])
- info = get_mla_metadata_info_v1(bs, q_len, NUM_HEADS, q_fp8.dtype, kv_fp8.dtype,
- is_sparse=False, fast_mode=False, num_kv_splits=num_splits, intra_batch_mode=True)
- work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
- kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
+ key = (bs, num_splits, str(q_fp8.dtype), str(kv_fp8.dtype))
+ if key not in _meta_cache:
+ info = get_mla_metadata_info_v1(bs, q_len, NUM_Q_HEADS, q_fp8.dtype, kv_fp8.dtype,
+ is_sparse=False, fast_mode=False, num_kv_splits=num_splits, intra_batch_mode=True)
+ _meta_cache[key] = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
+ work = _meta_cache[key]
kv_last = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
-
get_mla_metadata_v1(qo_indptr, kv_indptr, kv_last,
- NUM_HEADS, NUM_KV_HEADS, True,
+ NUM_Q_HEADS, NUM_KV_HEADS, True,
work[0], work[2], work[1], work[3], work[4], work[5],
page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16),
max_seqlen_qo=q_len, uni_seqlen_qo=q_len,
fast_mode=False, max_split_per_batch=num_splits,
intra_batch_mode=True, dtype_q=q_fp8.dtype, dtype_kv=kv_fp8.dtype)
- o = torch.empty((q.shape[0], NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
- mla_decode_fwd(q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM), kv_4d, o,
- qo_indptr, kv_indptr, kv_indices, kv_last, q_len,
+ if total_kv not in _idx_cache:
+ _idx_cache[total_kv] = torch.arange(total_kv, dtype=torch.int32, device="cuda")
+ if total_q not in _out_cache:
+ _out_cache[total_q] = torch.empty((total_q, NUM_Q_HEADS, V_DIM), dtype=torch.bfloat16, device="cuda")
+
+ o = _out_cache[total_q]
+ mla_decode_fwd(q_fp8.view(-1, NUM_Q_HEADS, QK_DIM), kv_4d, o,
+ qo_indptr, kv_indptr, _idx_cache[total_kv], kv_last, q_len,
page_size=PAGE_SIZE, nhead_kv=NUM_KV_HEADS, sm_scale=SM_SCALE, logit_cap=0.0,
num_kv_splits=num_splits, q_scale=q_scale, kv_scale=kv_scale,
intra_batch_mode=True,
⋯ 1 unchanged lines
reduce_indptr=work[3], reduce_final_map=work[4], reduce_partial_map=work[5])
return o
- def _sdpa_decode(q, kv_data, config):
- """Direct SDPA for small batches — bypasses flash attention overhead."""
- bs = config["batch_size"]
- kv_len = config["kv_seq_len"]
- # Use bf16 KV for SDPA (variable-length not supported, assume uniform kv_len)
- kv_bf16 = kv_data["bf16"] # (total_kv, 1, 576)
+ # ═══════════════════════════════════════════════
+ # BMM bf16 (for small shapes)
+ # ═══════════════════════════════════════════════
- # Reshape for batched attention
- Q = q.view(bs, 1, NUM_HEADS, QK_HEAD_DIM).transpose(1, 2) # (bs, 16, 1, 576)
- K = kv_bf16.view(bs, kv_len, 1, QK_HEAD_DIM).permute(0, 2, 1, 3).expand(bs, NUM_HEADS, kv_len, QK_HEAD_DIM) # (bs, 16, kv_len, 576)
- V = kv_bf16[:, :, :V_HEAD_DIM].view(bs, kv_len, 1, V_HEAD_DIM).permute(0, 2, 1, 3).expand(bs, NUM_HEADS, kv_len, V_HEAD_DIM) # (bs, 16, kv_len, 512)
-
- out = F.scaled_dot_product_attention(Q, K, V, scale=SM_SCALE, is_causal=False)
- return out.transpose(1, 2).reshape(bs, NUM_HEADS, V_HEAD_DIM) # (bs, 16, 512)
-
- def _bmm_decode(q, kv_data, config):
- """Direct bmm for smallest shapes — minimum overhead."""
+ def _bmm(q, kv_data, config):
bs = config["batch_size"]
kv_len = config["kv_seq_len"]
- kv_bf16 = kv_data["bf16"].view(bs, kv_len, QK_HEAD_DIM) # (bs, kv_len, 576)
+ kv = kv_data["bf16"].view(bs, kv_len, QK_DIM)
+ Q = q.view(bs, NUM_Q_HEADS, QK_DIM)
+ V = kv[:, :, :V_DIM]
+ s = torch.bmm(Q, kv.transpose(1, 2)) * SM_SCALE
+ w = F.softmax(s, dim=-1, dtype=torch.float32).to(torch.bfloat16)
+ return torch.bmm(w, V)
- Q = q.view(bs, NUM_HEADS, QK_HEAD_DIM) # (bs, 16, 576)
- K = kv_bf16 # (bs, kv_len, 576)
- V = kv_bf16[:, :, :V_HEAD_DIM] # (bs, kv_len, 512)
- # scores: (bs, 16, kv_len)
- scores = torch.bmm(Q, K.transpose(1, 2)) * SM_SCALE
- weights = F.softmax(scores, dim=-1).to(torch.bfloat16)
- # output: (bs, 16, 512)
- out = torch.bmm(weights, V)
- return out
+ # ═══════════════════════════════════════════════
+ # Shape dispatch based on measured data
+ # ═══════════════════════════════════════════════
+ # Per-shape best path (from benchmarks):
+ # bmm: bs=4/1k(26), bs=4/8k(43), bs=32/1k(42), bs=64/1k(61)
+ # Triton fp8: bs=32/8k(120) [AITER was 166]
+ # AITER fp8: bs=64/8k(206), bs=256/1k(161), bs=256/8k(353)
+
+ _AITER_SPLITS = {
+ (64, 8192): 32,
+ (256, 1024): 16,
+ (256, 8192): 32,
+ }
+
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
kv_len = config["kv_seq_len"]
+ sk = (bs, kv_len)
+ if sk not in _warm:
+ _warm.add(sk)
- # Shape dispatch — use fastest path per regime
- if bs <= 4:
- # Small batch: direct bmm is fastest (minimal overhead)
- return _bmm_decode(q, kv_data, config)
- else:
- # All other shapes: AITER fp8 (best throughput)
- return _aiter_fp8_decode(q, kv_data, qo_indptr, kv_indptr, config, num_splits=32)
+ # bmm: small shapes
+ if bs <= 4 or (bs <= 64 and kv_len <= 1024):
+ return _bmm(q, kv_data, config)
+
+ # Triton fp8: bs=32/kv=8192 (beats AITER)
+ if bs <= 32 and kv_len <= 8192:
+ return _triton_fp8(q, kv_data, kv_indptr, config, 8)
+
+ # AITER fp8: large shapes
+ return _aiter_fp8(q, kv_data, qo_indptr, kv_indptr, config, _AITER_SPLITS.get(sk, 32))
scrolls · 336 diff lines total

Best evidence level for this revision: reported

JSON