Skip to content
KernelIndex
Search⌘K

submission 668443

AlexNebula · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-668443?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
71.7µs
#334 of 766
2026-03-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:475a6001ab9635e1ea1f5ead6c6a3ce83b543543f5c3598b0a6a373e1f586cae
license declaredunknown
license concludedunknown
authorsAlexNebula
imported2026-08-26

Techniques

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

persistent-kernelNUM_KV_SPLITS = 32 # fixed — aiter persistent kernel needs enough work units
shared-memory__shared__ float wm[BLK/WSZ];

Kernel source

submission.py144 lines
"""
MLA Decode — hybrid a8w8/a16w8 with cached metadata.
v2: per-case NUM_KV_SPLITS tuning + fused fp8 quant via HIP.
"""
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ["CXX"] = "clang++"

import torch
from torch.utils.cpp_extension import load_inline
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_KV_HEADS = 1; QK_HEAD_DIM = 576; V_HEAD_DIM = 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5); PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8

# ===== HIP fast amax kernel (replace abs+amax = 2 kernel launches with 1) =====
_AMAX_CPP = """
#include <torch/extension.h>
torch::Tensor fast_amax(torch::Tensor input);
"""
_AMAX_HIP = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#define BLK 256
#define WSZ 64

__global__ void _amax_k(const __hip_bfloat16* __restrict__ in, float* __restrict__ out, int n) {
    float lm = 0.f;
    int gid = blockIdx.x * BLK + threadIdx.x;
    int stride = gridDim.x * BLK;
    // vectorized: process 4 bf16 at a time
    for (int i = gid * 4; i < n; i += stride * 4) {
        if (i + 3 < n) {
            uint64_t v = *(const uint64_t*)(in + i);
            __hip_bfloat16* bv = (__hip_bfloat16*)&v;
            lm = fmaxf(lm, fabsf(__bfloat162float(bv[0])));
            lm = fmaxf(lm, fabsf(__bfloat162float(bv[1])));
            lm = fmaxf(lm, fabsf(__bfloat162float(bv[2])));
            lm = fmaxf(lm, fabsf(__bfloat162float(bv[3])));
        } else {
            for (int j = i; j < min(i+4, n); j++)
                lm = fmaxf(lm, fabsf(__bfloat162float(in[j])));
        }
    }
    for (int o = WSZ/2; o > 0; o >>= 1)
        lm = fmaxf(lm, __shfl_down(lm, o));
    __shared__ float wm[BLK/WSZ];
    int lane = threadIdx.x % WSZ, wid = threadIdx.x / WSZ;
    if (lane == 0) wm[wid] = lm;
    __syncthreads();
    if (threadIdx.x == 0) {
        float m = wm[0];
        for (int i = 1; i < BLK/WSZ; i++) m = fmaxf(m, wm[i]);
        atomicMax((int*)out, __float_as_int(m));
    }
}

torch::Tensor fast_amax(torch::Tensor input) {
    int n = input.numel();
    int blocks = min(128, (n / 4 + BLK - 1) / BLK);
    blocks = max(1, blocks);
    auto result = torch::zeros({1}, torch::dtype(torch::kFloat32).device(input.device()));
    _amax_k<<<blocks, BLK>>>((const __hip_bfloat16*)input.data_ptr(),
        result.data_ptr<float>(), n);
    return result;
}
"""

try:
    _amax_mod = load_inline(name='famax', cpp_sources=[_AMAX_CPP], cuda_sources=[_AMAX_HIP],
        functions=['fast_amax'], verbose=True,
        extra_cuda_cflags=["--offload-arch=gfx950","-std=c++20","-O3"])
    _HAS_HIP_AMAX = True
except Exception:
    _HAS_HIP_AMAX = False

NUM_KV_SPLITS = 32  # fixed — aiter persistent kernel needs enough work units

def _use_a8w8(bs, kvl):
    return bs >= 64 and kvl >= 8192

_cache = {}

def _ensure_cached(bs, kvl, qoi, kvi):
    fp8q = _use_a8w8(bs, kvl)
    nsplits = NUM_KV_SPLITS
    key = (bs, kvl, fp8q)
    if key in _cache:
        return _cache[key]
    tkv = bs * kvl
    ki = torch.arange(tkv, dtype=torch.int32, device="cuda")
    klp = torch.full((bs,), kvl, dtype=torch.int32, device="cuda")
    out = torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
    qd = FP8_DTYPE if fp8q else torch.bfloat16
    info = get_mla_metadata_info_v1(bs, 1, NUM_HEADS, qd, FP8_DTYPE,
        is_sparse=False, fast_mode=False, num_kv_splits=nsplits, intra_batch_mode=True)
    wk = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
    get_mla_metadata_v1(qoi, kvi, klp, NUM_HEADS // NUM_KV_HEADS, NUM_KV_HEADS, True,
        wk[0], wk[2], wk[1], wk[3], wk[4], wk[5],
        page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16),
        max_seqlen_qo=1, uni_seqlen_qo=1, fast_mode=False,
        max_split_per_batch=nsplits, intra_batch_mode=True,
        dtype_q=qd, dtype_kv=FP8_DTYPE)
    # Pre-alloc fp8 buffers for a8w8 path
    q_fp8 = torch.empty((bs, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda") if fp8q else None
    q_scale = torch.empty(1, dtype=torch.float32, device="cuda") if fp8q else None
    _cache[key] = {"ki": ki, "klp": klp, "out": out, "fp8q": fp8q, "nsplits": nsplits,
        "q_fp8": q_fp8, "q_scale": q_scale,
        "meta": {"work_meta_data": wk[0], "work_indptr": wk[1], "work_info_set": wk[2],
                 "reduce_indptr": wk[3], "reduce_final_map": wk[4], "reduce_partial_map": wk[5]}}
    return _cache[key]

def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qoi, kvi, cfg = data
    bs = cfg["batch_size"]; kvl = cfg["kv_seq_len"]
    c = _ensure_cached(bs, kvl, qoi, kvi)
    kv8, kvs = kv_data["fp8"]
    kv4d = kv8.view(kv8.shape[0], PAGE_SIZE, NUM_KV_HEADS, -1)
    o = c["out"]

    if c["fp8q"]:
        fi = torch.finfo(FP8_DTYPE)
        if _HAS_HIP_AMAX:
            # HIP amax: 1 kernel instead of abs()+amax() = 2 kernels
            am = _amax_mod.fast_amax(q).clamp(min=1e-12)
        else:
            am = q.abs().amax().clamp(min=1e-12)
        sc = am.div_(fi.max)
        qi = q.div(sc).clamp_(min=fi.min, max=fi.max).to(FP8_DTYPE)
        qs = sc.float().reshape(1)
    else:
        qi, qs = q, None

    mla_decode_fwd(qi.view(-1, NUM_HEADS, QK_HEAD_DIM), kv4d, o, qoi, kvi,
        c["ki"], c["klp"], 1, page_size=PAGE_SIZE, nhead_kv=NUM_KV_HEADS,
        sm_scale=SM_SCALE, logit_cap=0.0, num_kv_splits=c["nsplits"],
        q_scale=qs, kv_scale=kvs, intra_batch_mode=True, **c["meta"])
    return o
scrolls · 144 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