Skip to content
KernelIndex
Search⌘K

submission 739756

divc13 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v94.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-739756?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
37.2µs
#91 of 766
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:7186d9ddb65c333c84f7bef41bd0cf2d2e00b9f324593c16154e7edcbe9318ee
license declaredunknown
license concludedunknown
authorsdivc13
imported2026-08-15

Techniques

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

autotune- bs32/64 kv1k: low splits + ns=3 (v92 autotune), less stage2 overhead
fp8FP8_DTYPE = tl.float8e4nv
mmaqk = tl.dot(q_nope, k_nope) + tl.dot(q_rope, k_rope)
num-warps = 4num_warps=4, num_stages=ns, waves_per_eu=2,
stages = 2num_warps=4, num_stages=2,
tile-n = 64NUM_KV_SPLITS=num_splits, BLOCK_N=64, BLOCK_H=16,

Kernel source

submission_v94.py233 lines
"""MLA decode Triton v94: cherry-picked optimal (splits, ns) per shape.

Best configs from v84/v92/v93 experiments:
- bs4: high splits (v84), GPU underutilized
- bs32/64 kv1k: low splits + ns=3 (v92 autotune), less stage2 overhead
- bs32/64 kv8k: moderate splits (v93 autotune)
- bs256: force_single kv1k, 2-split kv8k (v93 autotune)
"""

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

_SHAPE_CONFIGS = {
    (4, 1024):   (16, 2),
    (4, 8192):   (32, 3),
    (32, 1024):  (4,  3),
    (32, 8192):  (8,  3),
    (64, 1024):  (4,  3),
    (64, 8192):  (8,  2),
    (256, 1024): (1,  2),
    (256, 8192): (2,  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 _stage1_kernel(
    Q, KV, Mid_O, Out,
    qo_indptr, kv_indptr,
    sm_scale, kv_scale_ptr,
    stride_qt, stride_qh, stride_kvt,
    stride_mb, stride_mh, stride_ms,
    stride_ot, stride_oh,
    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,
):
    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.exp(e_max - new_max)
        p = tl.exp(qk - new_max[:, None])
        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.exp(e_max - new_max)
        p = tl.exp(qk - new_max[:, None])
        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)
    result = acc / safe_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:
        mid_base = cur_batch * stride_mb + heads * stride_mh + split_id * stride_ms
        tl.store(Mid_O + mid_base[:, None] + offs_c[None, :], result)
        lse = tl.where(e_sum > 0, e_max + tl.log(e_sum), float("-inf"))
        tl.store(Mid_O + mid_base + 512, lse)


@triton.jit
def _stage2_kernel(
    Mid_O, Out, qo_indptr, kv_scale_ptr,
    stride_mb, stride_mh, stride_ms, stride_ot, stride_oh,
    NUM_KV_SPLITS: tl.constexpr, BLOCK_DV: tl.constexpr,
    batch: tl.constexpr, GRID_SIZE: 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)
    mid_base = cur_batch * stride_mb + cur_head * stride_mh
    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)
        new_max = tl.maximum(lse, e_max)
        old_scale = tl.exp(e_max - new_max)
        exp_lse = tl.exp(lse - new_max)
        acc = acc * old_scale + exp_lse * tv
        e_sum = e_sum * old_scale + exp_lse
        e_max = new_max
    result = acc / tl.maximum(e_sum, 1e-12) * 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_config(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:
        splits = min(2, max_tiles)
    elif batch_size >= 64:
        splits = 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
    grid_size = batch_size * splits
    ns = 3 if grid_size < 512 else 2
    return splits, ns


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, ns = _get_config(batch_size, kv_len)

    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)

    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,
        sm_scale, kv_scale_t,
        q.stride(0), q.stride(1), kv.stride(0),
        mid_o.stride(0) if not direct_out else 0,
        mid_o.stride(1) if not direct_out else 0,
        mid_o.stride(2) if not direct_out else 0,
        out.stride(0), out.stride(1),
        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,
        num_warps=4, num_stages=ns, 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_o.stride(0), mid_o.stride(1), mid_o.stride(2),
            out.stride(0), out.stride(1),
            NUM_KV_SPLITS=num_splits, BLOCK_DV=512,
            batch=batch_size, GRID_SIZE=grid_size2,
            num_warps=4, num_stages=2,
        )

    return out
scrolls · 233 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 710716.

+ """MLA decode Triton v94: cherry-picked optimal (splits, ns) per shape.
+
+ Best configs from v84/v92/v93 experiments:
+ - bs4: high splits (v84), GPU underutilized
+ - bs32/64 kv1k: low splits + ns=3 (v92 autotune), less stage2 overhead
+ - bs32/64 kv8k: moderate splits (v93 autotune)
+ - bs256: force_single kv1k, 2-split kv8k (v93 autotune)
+ """
+
import torch
- import os
- from torch.utils.cpp_extension import load_inline
+ import triton
+ import triton.language as tl
from task import input_t, output_t
- # ---------------------------------------------------------------------------
- # MLA decode v249: Tuned split strategy + parallel pack
- # - Target 912 grid for large workloads, 768+ for medium
- # - Parallel packing: 2-step lshl_or chain instead of 3-step
- # - Base: v245 (restored v232 barriers + ubfe/lshl_or)
- # ---------------------------------------------------------------------------
+ NUM_HEADS = 16
+ KV_LORA_RANK = 512
+ QK_HEAD_DIM = 576
+ V_HEAD_DIM = 512
+ FP8_DTYPE = tl.float8e4nv
- os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
-
- HIP_SRC = r"""
- #include <torch/extension.h>
- #include <hip/hip_runtime.h>
-
- static constexpr int WARP_SIZE = 64;
- static constexpr int NUM_WARPS = 4;
- static constexpr int BLOCK_SIZE = WARP_SIZE * NUM_WARPS;
-
- static constexpr int QK_DIM = 576;
- static constexpr int V_DIM = 512;
- static constexpr int NUM_HEADS = 16;
- static constexpr int NUM_K_CHUNKS = QK_DIM / 32; // 18
- static constexpr int NUM_K128_CHUNKS = (QK_DIM + 127) / 128; // 5
- static constexpr int SUPER_TILE = 32;
- static constexpr int SV_CHUNKS = 8;
- static constexpr int KV_TILE_BYTES = SUPER_TILE * QK_DIM; // 18432
-
- // 18432 bytes / 16 bytes per uint4 / 256 threads = 4.5 -> 5 rounds
- static constexpr int PF_UINT4S = KV_TILE_BYTES / 16; // 1152
- static constexpr int PF_ROUNDS = (PF_UINT4S + BLOCK_SIZE - 1) / BLOCK_SIZE; // 5
-
- typedef float __attribute__((ext_vector_type(4))) v4f32;
- typedef unsigned int __attribute__((ext_vector_type(4))) u32x4;
- typedef int __attribute__((ext_vector_type(4))) i32x4;
- typedef int __attribute__((ext_vector_type(8))) i32x8;
- typedef unsigned int __attribute__((address_space(3)))* lds_ptr_t;
-
- extern "C" __device__ void __llvm_amdgcn_raw_buffer_load_lds(
- i32x4 rsrc, lds_ptr_t lds_ptr, int size,
- int voffset, int soffset, int offset, int aux)
- __asm("llvm.amdgcn.raw.buffer.load.lds");
-
- struct buffer_resource { uint64_t ptr; uint32_t range; uint32_t config; };
-
- __device__ __forceinline__ i32x4 make_buffer_rsrc(const void* p, uint32_t bytes) {
- buffer_resource r = {reinterpret_cast<uint64_t>(p), bytes, 0x110000};
- return *reinterpret_cast<i32x4*>(&r);
+ _SHAPE_CONFIGS = {
+ (4, 1024): (16, 2),
+ (4, 8192): (32, 3),
+ (32, 1024): (4, 3),
+ (32, 8192): (8, 3),
+ (64, 1024): (4, 3),
+ (64, 8192): (8, 2),
+ (256, 1024): (1, 2),
+ (256, 8192): (2, 2),
}
- __device__ __forceinline__ float bf16_to_f32(unsigned short v) {
- return __uint_as_float(static_cast<unsigned int>(v) << 16);
- }
- __device__ __forceinline__ unsigned short f32_to_bf16(float v) {
- unsigned int bits = __float_as_uint(v);
- bits += 0x7FFF + ((bits >> 16) & 1);
- return static_cast<unsigned short>(bits >> 16);
- }
+ @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
- __device__ __forceinline__ float fp8_to_f32(unsigned char b) {
- return __builtin_amdgcn_cvt_f32_fp8(static_cast<int>(b), 0);
- }
- __device__ __forceinline__ v4f32 mfma_f32_16x16x128_fp8(
- i32x8 A, i32x8 B, v4f32 C)
- {
- v4f32 D;
- asm(
- "v_mfma_f32_16x16x128_f8f6f4 %0, %1, %2, %3 cbsz:0 blgp:0"
- : "=v"(D) : "v"(A), "v"(B), "v"(C));
- return D;
- }
+ @triton.jit
+ def _stage1_kernel(
+ Q, KV, Mid_O, Out,
+ qo_indptr, kv_indptr,
+ sm_scale, kv_scale_ptr,
+ stride_qt, stride_qh, stride_kvt,
+ stride_mb, stride_mh, stride_ms,
+ stride_ot, stride_oh,
+ 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,
+ ):
+ 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
- __device__ __forceinline__ unsigned int lshl_or(unsigned int src, unsigned int shift, unsigned int base) {
- unsigned int r;
- asm("v_lshl_or_b32 %0, %1, %2, %3" : "=v"(r) : "v"(src), "v"(shift), "v"(base));
- return r;
- }
+ 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.exp(e_max - new_max)
+ p = tl.exp(qk - new_max[:, None])
+ 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
- __device__ __forceinline__ unsigned int pack4_fp8(unsigned int b0, unsigned int b1,
- unsigned int b2, unsigned int b3) {
- unsigned int lo = lshl_or(b1, 8u, b0);
- unsigned int hi = lshl_or(b3, 8u, b2);
- return lshl_or(hi, 16u, lo);
- }
+ 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.exp(e_max - new_max)
+ p = tl.exp(qk - new_max[:, None])
+ 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
- // =========================================================================
- // Tile processing: template specialization for full vs partial tiles
- // FULL_TILE=true: stcnt=32 (compile-time), all bounds checks eliminated
- // FULL_TILE=false: runtime stcnt with full bounds checks
- // =========================================================================
+ safe_esum = tl.where(e_sum > 0, e_sum, 1.0)
+ result = acc / safe_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:
+ mid_base = cur_batch * stride_mb + heads * stride_mh + split_id * stride_ms
+ tl.store(Mid_O + mid_base[:, None] + offs_c[None, :], result)
+ lse = tl.where(e_sum > 0, e_max + tl.log(e_sum), float("-inf"))
+ tl.store(Mid_O + mid_base + 512, lse)
- template<bool FULL_TILE>
- __device__ __forceinline__ void process_tile(
- const unsigned char* __restrict__ kv_ptr,
- unsigned char kv_lds[][KV_TILE_BYTES],
- float s_W[][16][33],
- const i32x8* q_128,
- const float score_scale,
- float mv[4], float lv[4],
- float vacc[][4],
- int& cur_buf,
- const int mr, const int kg, const int tid, const int warp_id,
- const int split_kv_start, const int split_kv_end,
- const int total_tokens, const int num_st,
- const int st_idx, const int stcnt_arg)
- {
- const int stcnt = FULL_TILE ? SUPER_TILE : stcnt_arg;
- const int ta = FULL_TILE ? 16 : min(16, stcnt);
- const int tb = FULL_TILE ? 16 : max(0, stcnt - 16);
- const unsigned char* kv_cur = kv_lds[cur_buf];
- // ---- INTERLEAVED DMA + QK: 1 DMA round per QK chunk ----
- // PF_ROUNDS == NUM_K128_CHUNKS == 5, so natural 1:1 interleaving.
- // DMA writes to kv_lds[nxt_buf] via VMEM, QK reads kv_lds[cur_buf] via LDS.
- __builtin_amdgcn_s_setprio(3);
- const int nxt_buf = cur_buf ^ 1;
- const bool has_next = (st_idx + 1 < num_st);
+ @triton.jit
+ def _stage2_kernel(
+ Mid_O, Out, qo_indptr, kv_scale_ptr,
+ stride_mb, stride_mh, stride_ms, stride_ot, stride_oh,
+ NUM_KV_SPLITS: tl.constexpr, BLOCK_DV: tl.constexpr,
+ batch: tl.constexpr, GRID_SIZE: 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)
+ mid_base = cur_batch * stride_mb + cur_head * stride_mh
+ 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)
+ new_max = tl.maximum(lse, e_max)
+ old_scale = tl.exp(e_max - new_max)
+ exp_lse = tl.exp(lse - new_max)
+ acc = acc * old_scale + exp_lse * tv
+ e_sum = e_sum * old_scale + exp_lse
+ e_max = new_max
+ result = acc / tl.maximum(e_sum, 1e-12) * kv_scale
+ o_base = q_tok * stride_ot + cur_head * stride_oh
+ tl.store(Out + o_base + offs_d, result.to(tl.bfloat16))
- i32x4 srsrc = {};
- if (has_next) {
- const int nxt_start = split_kv_start + (st_idx + 1) * SUPER_TILE;
- const int nxt_bytes = min(SUPER_TILE, split_kv_end - nxt_start) * QK_DIM;
- const unsigned char* __restrict__ nsrc = kv_ptr +
- static_cast<long long>(nxt_start) * QK_DIM;
- srsrc = make_buffer_rsrc(nsrc, nxt_bytes);
- }
- v4f32 ca = {0, 0, 0, 0};
- v4f32 cb = {0, 0, 0, 0};
- #pragma unroll
- for (int c = 0; c < NUM_K128_CHUNKS; c++) {
- if (has_next) {
- int dma_off = tid * 16 + c * BLOCK_SIZE * 16;
- lds_ptr_t ldp = (lds_ptr_t)(reinterpret_cast<uintptr_t>(kv_lds[nxt_buf]) + dma_off);
- __llvm_amdgcn_raw_buffer_load_lds(srsrc, ldp, 16, dma_off, 0, 0, 2);
- }
- i32x8 ba = {};
- if (FULL_TILE || mr < ta) {
- int base1 = mr * QK_DIM + c * 128 + 16 * kg;
- #pragma unroll
- for (int i = 0; i < 4; i++) {
- int off = base1 + i * 4;
- if (c * 128 + 16 * kg + i * 4 + 4 <= QK_DIM)
- ba[i] = *reinterpret_cast<const int*>(&kv_cur[off]);
- }
- int base2 = mr * QK_DIM + c * 128 + 64 + 16 * kg;
- #pragma unroll
- for (int i = 0; i < 4; i++) {
- int off = base2 + i * 4;
- if (c * 128 + 64 + 16 * kg + i * 4 + 4 <= QK_DIM)
- ba[4 + i] = *reinterpret_cast<const int*>(&kv_cur[off]);
- }
- }
- i32x8 bb = {};
- if (FULL_TILE || mr < tb) {
- int base1 = (16 + mr) * QK_DIM + c * 128 + 16 * kg;
- #pragma unroll
- for (int i = 0; i < 4; i++) {
- int off = base1 + i * 4;
- if (c * 128 + 16 * kg + i * 4 + 4 <= QK_DIM)
- bb[i] = *reinterpret_cast<const int*>(&kv_cur[off]);
- }
- int base2 = (16 + mr) * QK_DIM + c * 128 + 64 + 16 * kg;
- #pragma unroll
- for (int i = 0; i < 4; i++) {
- int off = base2 + i * 4;
- if (c * 128 + 64 + 16 * kg + i * 4 + 4 <= QK_DIM)
- bb[4 + i] = *reinterpret_cast<const int*>(&kv_cur[off]);
- }
- }
- ca = mfma_f32_16x16x128_fp8(q_128[c], ba, ca);
- cb = mfma_f32_16x16x128_fp8(q_128[c], bb, cb);
- }
- __builtin_amdgcn_s_setprio(0);
-
- // ---- Online softmax (16-lane reduce only: offsets 8,4,2,1) ----
- float sa0 = ca[0] * score_scale, sa1 = ca[1] * score_scale;
- float sa2 = ca[2] * score_scale, sa3 = ca[3] * score_scale;
- float sb0 = cb[0] * score_scale, sb1 = cb[1] * score_scale;
- float sb2 = cb[2] * score_scale, sb3 = cb[3] * score_scale;
-
- if (!FULL_TILE && mr >= ta) { sa0 = sa1 = sa2 = sa3 = -1e30f; }
- if (!FULL_TILE && mr >= tb) { sb0 = sb1 = sb2 = sb3 = -1e30f; }
-
- float tm0 = fmaxf(sa0, sb0), tm1 = fmaxf(sa1, sb1);
- float tm2 = fmaxf(sa2, sb2), tm3 = fmaxf(sa3, sb3);
- #pragma unroll
- for (int off = 8; off >= 1; off >>= 1) {
- tm0 = fmaxf(tm0, __shfl_xor(tm0, off));
- tm1 = fmaxf(tm1, __shfl_xor(tm1, off));
- tm2 = fmaxf(tm2, __shfl_xor(tm2, off));
- tm3 = fmaxf(tm3, __shfl_xor(tm3, off));
- }
-
- float nm0 = fmaxf(mv[0], tm0), nm1 = fmaxf(mv[1], tm1);
- float nm2 = fmaxf(mv[2], tm2), nm3 = fmaxf(mv[3], tm3);
- float rc0 = __expf(mv[0] - nm0), rc1 = __expf(mv[1] - nm1);
- float rc2 = __expf(mv[2] - nm2), rc3 = __expf(mv[3] - nm3);
- mv[0] = nm0; mv[1] = nm1; mv[2] = nm2; mv[3] = nm3;
-
- float wa0 = (FULL_TILE || mr < ta) ? __expf(sa0 - nm0) : 0.f;
- float wa1 = (FULL_TILE || mr < ta) ? __expf(sa1 - nm1) : 0.f;
- float wa2 = (FULL_TILE || mr < ta) ? __expf(sa2 - nm2) : 0.f;
- float wa3 = (FULL_TILE || mr < ta) ? __expf(sa3 - nm3) : 0.f;
- float wb0 = (FULL_TILE || mr < tb) ? __expf(sb0 - nm0) : 0.f;
- float wb1 = (FULL_TILE || mr < tb) ? __expf(sb1 - nm1) : 0.f;
- float wb2 = (FULL_TILE || mr < tb) ? __expf(sb2 - nm2) : 0.f;
- float wb3 = (FULL_TILE || mr < tb) ? __expf(sb3 - nm3) : 0.f;
-
- float dl0 = wa0 + wb0, dl1 = wa1 + wb1;
- float dl2 = wa2 + wb2, dl3 = wa3 + wb3;
- #pragma unroll
- for (int off = 8; off >= 1; off >>= 1) {
- dl0 += __shfl_xor(dl0, off); dl1 += __shfl_xor(dl1, off);
- dl2 += __shfl_xor(dl2, off); dl3 += __shfl_xor(dl3, off);
- }
- lv[0] = lv[0] * rc0 + dl0; lv[1] = lv[1] * rc1 + dl1;
- lv[2] = lv[2] * rc2 + dl2; lv[3] = lv[3] * rc3 + dl3;
-
- // ---- W to per-warp LDS ----
- s_W[warp_id][kg * 4 ][mr] = wa0;
- s_W[warp_id][kg * 4 + 1][mr] = wa1;
- s_W[warp_id][kg * 4 + 2][mr] = wa2;
- s_W[warp_id][kg * 4 + 3][mr] = wa3;
- s_W[warp_id][kg * 4 ][16 + mr] = wb0;
- s_W[warp_id][kg * 4 + 1][16 + mr] = wb1;
- s_W[warp_id][kg * 4 + 2][16 + mr] = wb2;
- s_W[warp_id][kg * 4 + 3][16 + mr] = wb3;
-
- asm volatile("s_waitcnt lgkmcnt(0)" ::: "memory");
-
- float wvals[8];
- #pragma unroll
- for (int i = 0; i < 8; i++)
- wvals[i] = s_W[warp_id][mr][i * 4 + kg];
-
- unsigned int wlo = __builtin_amdgcn_cvt_pk_fp8_f32(wvals[0], wvals[1], 0, false);
- wlo = __builtin_amdgcn_cvt_pk_fp8_f32(wvals[2], wvals[3], wlo, true);
- unsigned int whi = __builtin_amdgcn_cvt_pk_fp8_f32(wvals[4], wvals[5], 0, false);
- whi = __builtin_amdgcn_cvt_pk_fp8_f32(wvals[6], wvals[7], whi, true);
- long w_a = static_cast<long>(wlo) | (static_cast<long>(whi) << 32);
-
- // ---- SV MFMA from current LDS (each warp: 128 V dims) ----
- __builtin_amdgcn_s_setprio(1);
- const int v_warp_base = warp_id * SV_CHUNKS * 16;
-
- #pragma unroll
- for (int vc = 0; vc < SV_CHUNKS; vc += 2) {
- const int vd0 = v_warp_base + vc * 16 + mr;
- const int vd0_align = vd0 & ~3;
- const int vd0_shift = (vd0 & 3) * 8;
- const int vd1 = v_warp_base + (vc + 1) * 16 + mr;
- const int vd1_align = vd1 & ~3;
- const int vd1_shift = (vd1 & 3) * 8;
-
- unsigned int blo0 = 0, bhi0 = 0;
- unsigned int blo1 = 0, bhi1 = 0;
- if (vd0 < V_DIM) {
- unsigned int d0[8], d1[8];
- #pragma unroll
- for (int i = 0; i < 8; i++) {
- int tok = i * 4 + kg;
- if (FULL_TILE || tok < stcnt) {
- const unsigned int* base = reinterpret_cast<const unsigned int*>(
- &kv_cur[tok * QK_DIM + vd0_align]);
- d0[i] = base[0];
- d1[i] = base[4];
- } else {
- d0[i] = 0u;
- d1[i] = 0u;
- }
- }
- __builtin_amdgcn_sched_group_barrier(0x0040, 16, 0);
- __builtin_amdgcn_sched_group_barrier(0x0002, 64, 0);
- unsigned int e0[8], e1[8];
- #pragma unroll
- for (int i = 0; i < 8; i++) {
- e0[i] = __builtin_amdgcn_ubfe(d0[i], vd0_shift, 8);
- e1[i] = __builtin_amdgcn_ubfe(d1[i], vd1_shift, 8);
- }
- blo0 = pack4_fp8(e0[0], e0[1], e0[2], e0[3]);
- bhi0 = pack4_fp8(e0[4], e0[5], e0[6], e0[7]);
- blo1 = pack4_fp8(e1[0], e1[1], e1[2], e1[3]);
- bhi1 = pack4_fp8(e1[4], e1[5], e1[6], e1[7]);
- }
- long v_b0 = static_cast<long>(blo0) | (static_cast<long>(bhi0) << 32);
- long v_b1 = static_cast<long>(blo1) | (static_cast<long>(bhi1) << 32);
-
- v4f32 sc0 = {vacc[vc][0]*rc0, vacc[vc][1]*rc1, vacc[vc][2]*rc2, vacc[vc][3]*rc3};
- v4f32 sc1 = {vacc[vc+1][0]*rc0, vacc[vc+1][1]*rc1, vacc[vc+1][2]*rc2, vacc[vc+1][3]*rc3};
- sc0 = __builtin_amdgcn_mfma_f32_16x16x32_fp8_fp8(w_a, v_b0, sc0, 0, 0, 0);
- sc1 = __builtin_amdgcn_mfma_f32_16x16x32_fp8_fp8(w_a, v_b1, sc1, 0, 0, 0);
- vacc[vc][0] = sc0[0]; vacc[vc][1] = sc0[1];
- vacc[vc][2] = sc0[2]; vacc[vc][3] = sc0[3];
- vacc[vc+1][0] = sc1[0]; vacc[vc+1][1] = sc1[1];
- vacc[vc+1][2] = sc1[2]; vacc[vc+1][3] = sc1[3];
- }
-
- // ---- Wait for GLOBAL_LOAD_LDS and flip ----
- __builtin_amdgcn_s_setprio(3);
- if (has_next) {
- asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
- __builtin_amdgcn_sched_barrier(0);
- __builtin_amdgcn_s_barrier();
- cur_buf = nxt_buf;
- }
- }
-
- // =========================================================================
- // Full MFMA pipeline kernel with double-buffered LDS + K=128 QK MFMA
- // =========================================================================
-
- __global__ __launch_bounds__(256, 3)
- void mla_mfma_pipeline_kernel(
- const unsigned short* __restrict__ q_ptr,
- const unsigned char* __restrict__ kv_ptr,
- float* __restrict__ partial_m,
- float* __restrict__ partial_l,
- float* __restrict__ partial_acc,
- unsigned short* __restrict__ out_ptr,
- const int* __restrict__ qo_indptr,
- const int* __restrict__ kv_indptr,
- const float* __restrict__ kv_scale_ptr,
- const int num_splits,
- const float sm_scale)
- {
- const int split_idx = blockIdx.x;
- const int batch_idx = blockIdx.y;
- const int warp_id = threadIdx.x / WARP_SIZE;
- const int lane_id = threadIdx.x % WARP_SIZE;
- const int tid = threadIdx.x;
-
- const float score_scale = sm_scale * (*kv_scale_ptr);
-
- const int q_start = qo_indptr[batch_idx];
- const int kv_start = kv_indptr[batch_idx];
- const int kv_end = kv_indptr[batch_idx + 1];
- const int kv_len = kv_end - kv_start;
-
- const int tps = (kv_len + num_splits - 1) / num_splits;
- const int split_kv_start = kv_start + split_idx * tps;
- const int split_kv_end = min(split_kv_start + tps, kv_end);
-
- const int mr = lane_id & 0xF;
- const int kg = lane_id >> 4;
-
- if (split_kv_start >= kv_end) {
- if (lane_id < 16 && warp_id == 0) {
- int head = lane_id;
- int off = (batch_idx * NUM_HEADS + head) * num_splits + split_idx;
- partial_m[off] = -1e30f;
- partial_l[off] = 0.0f;
- }
- return;
- }
-
- // ===== LDS: double-buffered KV + per-warp W =====
- __shared__ __align__(16) unsigned char kv_lds[2][KV_TILE_BYTES];
- __shared__ float s_W[NUM_WARPS][16][33];
-
- // ===== Super-tile iteration setup =====
- const int total_tokens = split_kv_end - split_kv_start;
- const int num_st = (total_tokens + SUPER_TILE - 1) / SUPER_TILE;
-
- // ===== PROLOGUE: issue DMA FIRST, then Q prep overlaps with DMA =====
- {
- const int first_bytes = min(SUPER_TILE, total_tokens) * QK_DIM;
- const unsigned char* __restrict__ src0 = kv_ptr +
- static_cast<long long>(split_kv_start) * QK_DIM;
- i32x4 srsrc = make_buffer_rsrc(src0, first_bytes);
- #pragma unroll
- for (int r = 0; r < PF_ROUNDS; r++) {
- int off = tid * 16 + r * BLOCK_SIZE * 16;
- lds_ptr_t ldp = (lds_ptr_t)(reinterpret_cast<uintptr_t>(kv_lds[0]) + off);
- __llvm_amdgcn_raw_buffer_load_lds(srsrc, ldp, 16, off, 0, 0, 2);
- }
- }
-
- // ===== Q preload as FP8 (overlapped with DMA in flight) =====
- const unsigned short* qh = q_ptr +
- (static_cast<long long>(q_start) * NUM_HEADS + mr) * QK_DIM;
-
- i32x8 q_128[NUM_K128_CHUNKS];
- #pragma unroll
- for (int c = 0; c < NUM_K128_CHUNKS; c++) {
- unsigned int w[8];
- int base1 = c * 128 + 16 * kg;
- #pragma unroll
- for (int i = 0; i < 4; i++) {
- int d = base1 + i * 4;
- float f0 = (d < QK_DIM) ? bf16_to_f32(qh[d]) : 0.f;
- float f1 = (d + 1 < QK_DIM) ? bf16_to_f32(qh[d + 1]) : 0.f;
- float f2 = (d + 2 < QK_DIM) ? bf16_to_f32(qh[d + 2]) : 0.f;
- float f3 = (d + 3 < QK_DIM) ? bf16_to_f32(qh[d + 3]) : 0.f;
- unsigned int pk = __builtin_amdgcn_cvt_pk_fp8_f32(f0, f1, 0, false);
- pk = __builtin_amdgcn_cvt_pk_fp8_f32(f2, f3, pk, true);
- w[i] = pk;
- }
- int base2 = c * 128 + 64 + 16 * kg;
- #pragma unroll
- for (int i = 0; i < 4; i++) {
- int d = base2 + i * 4;
- float f0 = (d < QK_DIM) ? bf16_to_f32(qh[d]) : 0.f;
- float f1 = (d + 1 < QK_DIM) ? bf16_to_f32(qh[d + 1]) : 0.f;
- float f2 = (d + 2 < QK_DIM) ? bf16_to_f32(qh[d + 2]) : 0.f;
- float f3 = (d + 3 < QK_DIM) ? bf16_to_f32(qh[d + 3]) : 0.f;
- unsigned int pk = __builtin_amdgcn_cvt_pk_fp8_f32(f0, f1, 0, false);
- pk = __builtin_amdgcn_cvt_pk_fp8_f32(f2, f3, pk, true);
- w[4 + i] = pk;
- }
- q_128[c] = *reinterpret_cast<i32x8*>(w);
- }
-
- // ===== V accumulators + softmax state =====
- float vacc[SV_CHUNKS][4];
- #pragma unroll
- for (int i = 0; i < SV_CHUNKS; i++)
- vacc[i][0] = vacc[i][1] = vacc[i][2] = vacc[i][3] = 0.0f;
-
- float mv[4] = {-1e30f, -1e30f, -1e30f, -1e30f};
- float lv[4] = {0.0f, 0.0f, 0.0f, 0.0f};
-
- // ===== Wait for DMA (Q prep ran while DMA was in flight) =====
- asm volatile("s_waitcnt vmcnt(0)" ::: "memory");
- __builtin_amdgcn_sched_barrier(0);
- __builtin_amdgcn_s_barrier();
-
- int cur_buf = 0;
-
- // ===== MAIN LOOP: Two-phase for compile-time optimization =====
- const int num_full = total_tokens / SUPER_TILE;
- const int has_partial = (total_tokens % SUPER_TILE) != 0;
-
- // Phase 1: Full tiles — stcnt=32 is compile-time constant
- for (int st_idx = 0; st_idx < num_full; st_idx++) {
- process_tile<true>(kv_ptr, kv_lds, s_W, q_128, score_scale,
- mv, lv, vacc, cur_buf, mr, kg, tid, warp_id,
- split_kv_start, split_kv_end, total_tokens, num_st,
- st_idx, SUPER_TILE);
- }
-
- // Phase 2: Last tile with runtime stcnt (if partial)
- if (has_partial) {
- int last_stcnt = total_tokens - num_full * SUPER_TILE;
- process_tile<false>(kv_ptr, kv_lds, s_W, q_128, score_scale,
- mv, lv, vacc, cur_buf, mr, kg, tid, warp_id,
- split_kv_start, split_kv_end, total_tokens, num_st,
- num_full, last_stcnt);
- }
-
- if (num_splits == 1) {
- const float kv_scale = *kv_scale_ptr;
- const int vwb = warp_id * SV_CHUNKS * 16;
- #pragma unroll
- for (int vc = 0; vc < SV_CHUNKS; vc++) {
- int vd = vwb + vc * 16 + mr;
- if (vd < V_DIM) {
- #pragma unroll
- for (int r = 0; r < 4; r++) {
- int head = kg * 4 + r;
- float inv_l = (lv[r] > 0.f) ? (kv_scale / lv[r]) : 0.f;
- long long idx = (static_cast<long long>(q_start) * NUM_HEADS + head) * V_DIM + vd;
- out_ptr[idx] = f32_to_bf16(vacc[vc][r] * inv_l);
- }
- }
- }
- } else {
- if (warp_id == 0 && mr == 0) {
- #pragma unroll
- for (int r = 0; r < 4; r++) {
- int head = kg * 4 + r;
- int off = (batch_idx * NUM_HEADS + head) * num_splits + split_idx;
- partial_m[off] = mv[r];
- partial_l[off] = lv[r];
- }
- }
- const int vwb = warp_id * SV_CHUNKS * 16;
- #pragma unroll
- for (int vc = 0; vc < SV_CHUNKS; vc++) {
- int vd = vwb + vc * 16 + mr;
- if (vd < V_DIM) {
- #pragma unroll
- for (int r = 0; r < 4; r++) {
- int head = kg * 4 + r;
- int off = (batch_idx * NUM_HEADS + head) * num_splits + split_idx;
- partial_acc[static_cast<long long>(off) * V_DIM + vd] = vacc[vc][r];
- }
- }
- }
- }
- }
-
- // =========================================================================
- // Reduce kernel (template-specialized for compile-time loop unrolling)
- // =========================================================================
-
- __global__ __launch_bounds__(512)
- void mla_reduce_kernel_generic(
- const float* __restrict__ partial_m,
- const float* __restrict__ partial_l,
- const float* __restrict__ partial_acc,
- unsigned short* __restrict__ out_ptr,
- const float* __restrict__ kv_scale_ptr,
- const int num_splits)
- {
- const int item_idx = blockIdx.x;
- const int tid = threadIdx.x;
- const float kv_scale = *kv_scale_ptr;
-
- __shared__ float s_corr[128];
- __shared__ float s_inv_l;
-
- const int base = item_idx * num_splits;
-
- float my_m = -1e30f;
- float my_l = 0.0f;
- if (tid < num_splits) {
- my_m = partial_m[base + tid];
- my_l = partial_l[base + tid];
- }
-
- float merged_m = my_m;
- #pragma unroll
- for (int off = 32; off >= 1; off >>= 1)
- merged_m = fmaxf(merged_m, __shfl_xor(merged_m, off));
-
- float my_c = 0.0f;
- if (tid < num_splits && my_l > 0.f)
- my_c = __expf(my_m - merged_m);
- float weighted_l = my_l * my_c;
-
- if (tid < num_splits)
- s_corr[tid] = my_c;
-
- float merged_l = weighted_l;
- #pragma unroll
- for (int off = 32; off >= 1; off >>= 1)
- merged_l += __shfl_xor(merged_l, off);
-
- if (tid == 0)
- s_inv_l = (merged_l > 0.f) ? (kv_scale / merged_l) : 0.f;
- __syncthreads();
-
- if (tid < V_DIM) {
- float val = 0.0f;
- for (int s = 0; s < num_splits; ++s) {
- val += partial_acc[(static_cast<long long>(base + s)) * V_DIM + tid]
- * s_corr[s];
- }
- out_ptr[static_cast<long long>(item_idx) * V_DIM + tid] = f32_to_bf16(val * s_inv_l);
- }
- }
-
- template<int NUM_SPLITS>
- __global__ __launch_bounds__(512)
- void mla_reduce_kernel(
- const float* __restrict__ partial_m,
- const float* __restrict__ partial_l,
- const float* __restrict__ partial_acc,
- unsigned short* __restrict__ out_ptr,
- const float* __restrict__ kv_scale_ptr)
- {
- const int item_idx = blockIdx.x;
- const int tid = threadIdx.x;
- const float kv_scale = *kv_scale_ptr;
-
- __shared__ float s_corr[NUM_SPLITS < 128 ? 128 : NUM_SPLITS];
- __shared__ float s_inv_l;
-
- const int base = item_idx * NUM_SPLITS;
-
- float my_m = -1e30f;
- float my_l = 0.0f;
- if (tid < NUM_SPLITS) {
- my_m = partial_m[base + tid];
- my_l = partial_l[base + tid];
- }
-
- float merged_m = my_m;
- #pragma unroll
- for (int off = 32; off >= 1; off >>= 1)
- merged_m = fmaxf(merged_m, __shfl_xor(merged_m, off));
-
- float my_c = 0.0f;
- if (tid < NUM_SPLITS && my_l > 0.f)
- my_c = __expf(my_m - merged_m);
- float weighted_l = my_l * my_c;
-
- if (tid < NUM_SPLITS)
- s_corr[tid] = my_c;
-
- float merged_l = weighted_l;
- #pragma unroll
- for (int off = 32; off >= 1; off >>= 1)
- merged_l += __shfl_xor(merged_l, off);
-
- if (tid == 0)
- s_inv_l = (merged_l > 0.f) ? (kv_scale / merged_l) : 0.f;
- __syncthreads();
-
- if (tid < V_DIM) {
- float val = 0.0f;
- #pragma unroll
- for (int s = 0; s < NUM_SPLITS; ++s) {
- val += partial_acc[(static_cast<long long>(base + s)) * V_DIM + tid]
- * s_corr[s];
- }
- out_ptr[static_cast<long long>(item_idx) * V_DIM + tid] = f32_to_bf16(val * s_inv_l);
- }
- }
-
- torch::Tensor mla_decode(
- torch::Tensor q, torch::Tensor kv_buffer,
- torch::Tensor qo_indptr, torch::Tensor kv_indptr,
- torch::Tensor kv_scale_tensor,
- int64_t num_heads, int64_t num_splits,
- float sm_scale,
- torch::Tensor partial_m, torch::Tensor partial_l,
- torch::Tensor partial_acc,
- torch::Tensor output)
- {
- const int batch_size = qo_indptr.size(0) - 1;
- const int num_items = batch_size * static_cast<int>(num_heads);
-
- dim3 grid1(static_cast<int>(num_splits), batch_size);
- dim3 block1(BLOCK_SIZE);
- mla_mfma_pipeline_kernel<<<grid1, block1>>>(
- reinterpret_cast<const unsigned short*>(q.data_ptr()),
- reinterpret_cast<const unsigned char*>(kv_buffer.data_ptr()),
- partial_m.data_ptr<float>(), partial_l.data_ptr<float>(),
- partial_acc.data_ptr<float>(),
- reinterpret_cast<unsigned short*>(output.data_ptr()),
- qo_indptr.data_ptr<int>(), kv_indptr.data_ptr<int>(),
- kv_scale_tensor.data_ptr<float>(),
- static_cast<int>(num_splits), sm_scale);
- if (num_splits == 1) return output;
-
- dim3 grid2(num_items);
- dim3 block2(512);
-
- #define REDUCE_DISPATCH(N) \
- mla_reduce_kernel<N><<<grid2, block2>>>( \
- partial_m.data_ptr<float>(), partial_l.data_ptr<float>(), \
- partial_acc.data_ptr<float>(), \
- reinterpret_cast<unsigned short*>(output.data_ptr()), \
- kv_scale_tensor.data_ptr<float>())
-
- switch (static_cast<int>(num_splits)) {
- case 3: REDUCE_DISPATCH(3); break;
- case 8: REDUCE_DISPATCH(8); break;
- case 12: REDUCE_DISPATCH(12); break;
- case 16: REDUCE_DISPATCH(16); break;
- case 24: REDUCE_DISPATCH(24); break;
- case 64: REDUCE_DISPATCH(64); break;
- default:
- mla_reduce_kernel_generic<<<grid2, block2>>>(
- partial_m.data_ptr<float>(), partial_l.data_ptr<float>(),
- partial_acc.data_ptr<float>(),
- reinterpret_cast<unsigned short*>(output.data_ptr()),
- kv_scale_tensor.data_ptr<float>(),
- static_cast<int>(num_splits));
- break;
- }
- #undef REDUCE_DISPATCH
-
- return output;
- }
- """
-
- CPP_DECL = """
- torch::Tensor mla_decode(
- torch::Tensor q, torch::Tensor kv_buffer,
- torch::Tensor qo_indptr, torch::Tensor kv_indptr,
- torch::Tensor kv_scale_tensor,
- int64_t num_heads, int64_t num_splits,
- float sm_scale,
- torch::Tensor partial_m, torch::Tensor partial_l,
- torch::Tensor partial_acc,
- torch::Tensor output);
- """
-
- _module = load_inline(
- name="mla_hip_v249_tuned_splits",
- cpp_sources=CPP_DECL,
- cuda_sources=HIP_SRC,
- functions=["mla_decode"],
- extra_cuda_cflags=[
- "-O3", "-std=c++17",
- "-ffast-math", "-funsafe-math-optimizations", "-ffp-contract=fast",
- "-fno-gpu-rdc",
- "-mllvm", "-amdgpu-early-inline-all=true",
- "-mllvm", "-amdgpu-function-calls=false",
- "-mllvm", "-amdgpu-max-memory-clause=64",
- "-mllvm", "-amdgpu-load-store-vectorizer",
- "-mllvm", "-amdgpu-early-ifcvt",
- "-mllvm", "-amdgpu-internalize-symbols",
- "-mllvm", "-amdgpu-scalarize-global-loads",
- "-mllvm", "-amdgpu-dpp-combine",
- "-mllvm", "-amdgpu-enable-pre-ra-optimizations",
- "-mllvm", "-amdgpu-promote-alloca-to-vector-limit=256",
- ],
- verbose=False,
- )
-
-
-
- def _choose_splits(batch_size, kv_len):
- tiles = max(1, kv_len // 32)
- max_useful = max(1, tiles // 2)
-
- if tiles > 64:
- ideal = max(1, -(-768 // batch_size))
+ def _get_config(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:
+ splits = min(2, max_tiles)
+ elif batch_size >= 64:
+ splits = min(max(1, 512 // batch_size), max_tiles)
else:
- target_wgs = max(512, batch_size * 8)
- ideal = max(1, target_wgs // batch_size)
-
- splits = max(1, min(ideal, max_useful, 64))
-
+ splits = min(max(1, 768 // batch_size), max_tiles)
while splits > 1 and batch_size * splits > 912:
splits -= 1
+ grid_size = batch_size * splits
+ ns = 3 if grid_size < 512 else 2
+ return splits, ns
- return splits
-
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
-
- kv_buffer_fp8, kv_scale = kv_data["fp8"]
- kv_buffer = kv_buffer_fp8.view(-1, 576)
-
+ 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"]
- num_heads = config["num_heads"]
sm_scale = config["sm_scale"]
- total_q = q.size(0)
- num_items = batch_size * num_heads
+ total_q = q.shape[0]
+ kv_len = kv.shape[0] // batch_size if batch_size > 0 else 0
- total_kv = kv_buffer.shape[0]
- kv_len = total_kv // batch_size
+ num_splits, ns = _get_config(batch_size, kv_len)
- num_splits = _choose_splits(batch_size, kv_len)
-
dev = q.device
- pm = torch.empty((num_items * num_splits,), dtype=torch.float32, device=dev)
- pl = torch.empty((num_items * num_splits,), dtype=torch.float32, device=dev)
- pa = torch.empty((num_items * num_splits, 512), dtype=torch.float32, device=dev)
- out = torch.empty((total_q, num_heads, 512), dtype=torch.bfloat16, device=dev)
+ 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)
- return _module.mla_decode(
- q, kv_buffer, qo_indptr, kv_indptr,
- kv_scale,
- num_heads, num_splits,
- sm_scale,
- pm, pl, pa, out)
+ 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,
+ sm_scale, kv_scale_t,
+ q.stride(0), q.stride(1), kv.stride(0),
+ mid_o.stride(0) if not direct_out else 0,
+ mid_o.stride(1) if not direct_out else 0,
+ mid_o.stride(2) if not direct_out else 0,
+ out.stride(0), out.stride(1),
+ 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,
+ num_warps=4, num_stages=ns, 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_o.stride(0), mid_o.stride(1), mid_o.stride(2),
+ out.stride(0), out.stride(1),
+ NUM_KV_SPLITS=num_splits, BLOCK_DV=512,
+ batch=batch_size, GRID_SIZE=grid_size2,
+ num_warps=4, num_stages=2,
+ )
+
+ return out
scrolls · 961 diff lines total

Best evidence level for this revision: reported

JSON