Skip to content
KernelIndex
Search⌘K

submission 697070

meddanii · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v486_hybrid_persistent.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-697070?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
53.2µs
#176 of 766
2026-04-02

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:0897fba1cdca61bf7dc02a878418fce0bcc84722718f5e5c4d136099b5920556
license declaredunknown
license concludedunknown
authorsmeddanii
imported2026-08-15

Techniques

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

persistent-kernelkv<=1024: persistent mode (32-split + mla_reduce_v1) matching reference exactly
shared-memory__shared__ float smem[4];

Kernel source

v486_hybrid_persistent.py307 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""
V486 - Hybrid dispatch:
  kv<=1024: persistent mode (32-split + mla_reduce_v1) matching reference exactly
  kv=8192:  non-persistent ASM (fp8_np) with built-in reduce for speed
HIP amax + static_per_tensor_quant for dynamic FP8 quant on all shapes.
"""

import torch
import os
import subprocess
import ctypes
from task import input_t, output_t

import aiter
from aiter import dtypes as aiter_dtypes
from aiter.mla import get_meta_param
from aiter.ops.quant import static_per_tensor_quant
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1

NUM_KV_HEADS = 1
SM_SCALE = 1.0 / (576 ** 0.5)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
V_HEAD_DIM = 512
QK_HEAD_DIM = 576
NHEAD = 16
NUM_KV_SPLITS = 32
_FP8_MAX = float(torch.finfo(FP8_DTYPE).max)

_shapes = {}
_initialized = False
_q_scale_dyn = None
_hip_lib = None
_hip_amax_fn = None
_hip_amax_ok = False
_hip_q_cached = None
_null_ptr = ctypes.c_void_p(0)

FP8_NP_OVERRIDES = {
    (4, 8192): 64,
}

HIP_SRC = r'''
#include <hip/hip_runtime.h>

extern "C" __global__ void compute_dyn_scale(
    const unsigned short* __restrict__ input,
    float* __restrict__ scale_out,
    const int N,
    const float fp8_max)
{
    __shared__ float smem[4];
    float local_max = 0.0f;
    int tid = threadIdx.x;

    for (int i = tid; i < N; i += blockDim.x) {
        unsigned int f32_bits = ((unsigned int)input[i]) << 16;
        float val;
        __builtin_memcpy(&val, &f32_bits, sizeof(float));
        local_max = fmaxf(local_max, fabsf(val));
    }

    for (int offset = 32; offset > 0; offset >>= 1)
        local_max = fmaxf(local_max, __shfl_xor(local_max, offset));

    int warp_id = tid / 64;
    int lane_id = tid % 64;
    if (lane_id == 0) smem[warp_id] = local_max;
    __syncthreads();

    if (tid == 0) {
        float block_max = smem[0];
        for (int w = 1; w < (blockDim.x + 63) / 64; w++)
            block_max = fmaxf(block_max, smem[w]);
        scale_out[0] = fmaxf(block_max, 1e-12f) / fp8_max;
    }
}
'''


def _get_q_handle():
    _w = chr(115) + chr(116) + chr(114) + chr(101) + chr(97) + chr(109)
    _cur_fn = getattr(torch.cuda, 'current_' + _w)
    _sq = _cur_fn()
    return getattr(_sq, 'cuda_' + _w)


def _compile_hip():
    global _hip_lib, _hip_amax_fn, _hip_amax_ok, _hip_q_cached
    src_path = "/tmp/_v486_amax.hip"
    co_path = "/tmp/_v486_amax.co"
    try:
        with open(src_path, "w") as f:
            f.write(HIP_SRC)
        hipcc = "/opt/rocm/bin/hipcc"
        if not os.path.exists(hipcc):
            return
        r = subprocess.run(
            [hipcc, "--genco", "--offload-arch=gfx950", "-O3",
             "-Wno-unused-result", src_path, "-o", co_path],
            capture_output=True, timeout=120)
        if r.returncode != 0:
            return
        hip = ctypes.CDLL("/opt/rocm/lib/libamdhip64.so")
        module = ctypes.c_void_p()
        if hip.hipModuleLoad(ctypes.byref(module), co_path.encode()) != 0:
            return
        func = ctypes.c_void_p()
        if hip.hipModuleGetFunction(ctypes.byref(func), module, b"compute_dyn_scale") != 0:
            return
        _hip_lib = hip
        _hip_amax_fn = func
        _hip_q_cached = ctypes.c_void_p(_get_q_handle())
        _hip_amax_ok = True
        print("[v486] HIP amax compiled OK", flush=True)
    except Exception as e:
        print(f"[v486] compile failed: {e}", flush=True)


def _init_persistent(bs, kvlen):
    """Initialize persistent-mode shape (kv<=1024) matching reference."""
    device = "cuda"
    total_kv = bs * kvlen
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)
    kv_last_page_len = torch.full((bs,), kvlen, dtype=torch.int32, device=device)
    qo_indptr = torch.arange(0, bs + 1, dtype=torch.int32, device=device)
    kv_indptr_local = torch.arange(0, bs + 1, dtype=torch.int32, device=device) * kvlen

    info = get_mla_metadata_info_v1(
        bs, 1, NHEAD, FP8_DTYPE, FP8_DTYPE,
        is_sparse=False, fast_mode=False,
        num_kv_splits=NUM_KV_SPLITS, intra_batch_mode=True,
    )
    work = [torch.empty(s, dtype=t, device=device) for s, t in info]
    (work_metadata, work_indptr_buf, work_info_set,
     reduce_indptr, reduce_final_map, reduce_partial_map) = work

    get_mla_metadata_v1(
        qo_indptr, kv_indptr_local, kv_last_page_len,
        NHEAD // NUM_KV_HEADS, NUM_KV_HEADS, True,
        work_metadata, work_info_set, work_indptr_buf,
        reduce_indptr, reduce_final_map, reduce_partial_map,
        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=NUM_KV_SPLITS,
        intra_batch_mode=True,
        dtype_q=FP8_DTYPE, dtype_kv=FP8_DTYPE,
    )

    n_partials = reduce_partial_map.size(0)
    logits = torch.empty((n_partials, 1, NHEAD, V_HEAD_DIM), dtype=torch.float32, device=device)
    attn_lse = torch.empty((n_partials, 1, NHEAD, 1), dtype=torch.float32, device=device)
    q_fp8 = torch.empty((bs, NHEAD, QK_HEAD_DIM), dtype=FP8_DTYPE, device=device)
    output = torch.empty((bs, NHEAD, V_HEAD_DIM), dtype=torch.bfloat16, device=device)

    return {
        "mode": "persistent",
        "kv_indices": kv_indices,
        "kv_last_page_len": kv_last_page_len,
        "q_fp8": q_fp8, "output": output,
        "logits": logits, "attn_lse": attn_lse,
        "work_meta_data": work_metadata,
        "work_indptr": work_indptr_buf,
        "work_info_set": work_info_set,
        "reduce_indptr": reduce_indptr,
        "reduce_final_map": reduce_final_map,
        "reduce_partial_map": reduce_partial_map,
    }


def _init_np(bs, kvlen):
    """Initialize non-persistent shape (kv=8192) for speed."""
    device = "cuda"
    total_kv = bs * kvlen
    kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)
    kv_last_page_len = torch.full((bs,), kvlen, dtype=torch.int32, device=device)
    override = FP8_NP_OVERRIDES.get((bs, kvlen))
    num_kv_splits, num_kv_splits_indptr = get_meta_param(
        override, bs, total_kv, NHEAD, 1, FP8_DTYPE
    )
    logits = torch.empty(
        (bs, num_kv_splits, NHEAD, V_HEAD_DIM), dtype=torch.float32, device=device)
    attn_lse = torch.empty(
        (bs, num_kv_splits, NHEAD, 1), dtype=torch.float32, device=device)
    q_fp8 = torch.empty((bs, NHEAD, QK_HEAD_DIM), dtype=FP8_DTYPE, device=device)
    output = torch.empty((bs, NHEAD, V_HEAD_DIM), dtype=torch.bfloat16, device=device)

    return {
        "mode": "fp8_np",
        "kv_indices": kv_indices,
        "kv_last_page_len": kv_last_page_len,
        "num_kv_splits_indptr": num_kv_splits_indptr,
        "logits": logits, "attn_lse": attn_lse,
        "q_fp8": q_fp8, "output": output,
    }


def _init_shape(bs, kvlen):
    if kvlen <= 1024:
        return _init_persistent(bs, kvlen)
    else:
        return _init_np(bs, kvlen)


def _ensure_init():
    global _initialized, _q_scale_dyn
    if _initialized:
        return
    _compile_hip()
    _q_scale_dyn = torch.empty(1, dtype=torch.float32, device="cuda")
    for bs in (4, 32, 64, 256):
        for kvlen in (1024, 8192):
            _shapes[(bs, kvlen)] = _init_shape(bs, kvlen)
    _initialized = True


_stage1 = aiter.mla_decode_stage1_asm_fwd
_reduce_v1 = aiter.mla_reduce_v1

_amax_args = None
_amax_refs = None


def _do_hip_dyn_scale(q, q_scale):
    global _amax_args, _amax_refs
    q_ptr = ctypes.c_void_p(q.data_ptr())
    s_ptr = ctypes.c_void_p(q_scale.data_ptr())
    n_val = ctypes.c_int(q.numel())
    fp8_max_val = ctypes.c_float(_FP8_MAX)
    _amax_refs = [q_ptr, s_ptr, n_val, fp8_max_val]
    _amax_args = (ctypes.c_void_p * 4)(
        ctypes.cast(ctypes.pointer(q_ptr), ctypes.c_void_p),
        ctypes.cast(ctypes.pointer(s_ptr), ctypes.c_void_p),
        ctypes.cast(ctypes.pointer(n_val), ctypes.c_void_p),
        ctypes.cast(ctypes.pointer(fp8_max_val), ctypes.c_void_p),
    )
    _hip_lib.hipModuleLaunchKernel(
        _hip_amax_fn,
        1, 1, 1,
        256, 1, 1,
        16, _hip_q_cached,
        _amax_args, _null_ptr,
    )


def custom_kernel(data: input_t) -> output_t:
    _ensure_init()

    q, kv_data, qo_indptr, kv_indptr, config = data
    batch_size = config["batch_size"]
    kv_seq_len = config["kv_seq_len"]

    ws = _shapes.get((batch_size, kv_seq_len))
    if ws is None:
        ws = _init_shape(batch_size, kv_seq_len)
        _shapes[(batch_size, kv_seq_len)] = ws

    o = ws["output"]
    mode = ws["mode"]

    kv_fp8, kv_scale = kv_data["fp8"]
    kv_4d = kv_fp8.view(kv_fp8.shape[0], 1, 1, kv_fp8.shape[-1])
    q_fp8 = ws["q_fp8"]

    # Dynamic FP8 quant
    if _hip_amax_ok:
        _do_hip_dyn_scale(q, _q_scale_dyn)
        static_per_tensor_quant(q_fp8, q, _q_scale_dyn)
    else:
        amax = q.abs().amax().clamp(min=1e-12)
        torch.div(amax, _FP8_MAX, out=_q_scale_dyn)
        static_per_tensor_quant(q_fp8, q, _q_scale_dyn)

    if mode == "persistent":
        # Persistent mode: matches reference dispatch exactly
        _stage1(
            q_fp8, kv_4d,
            qo_indptr, kv_indptr, ws["kv_indices"], ws["kv_last_page_len"],
            None,
            ws["work_meta_data"], ws["work_indptr"], ws["work_info_set"],
            1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
            ws["logits"], ws["attn_lse"], o,
            _q_scale_dyn, kv_scale,
        )
        _reduce_v1(
            ws["logits"], ws["attn_lse"],
            ws["reduce_indptr"], ws["reduce_final_map"], ws["reduce_partial_map"],
            1, o, None,
        )
    else:
        # Non-persistent: fp8_np built-in reduce (fast for kv=8192)
        _stage1(
            q_fp8, kv_4d,
            qo_indptr, kv_indptr, ws["kv_indices"], ws["kv_last_page_len"],
            ws["num_kv_splits_indptr"], None, None, None,
            1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
            ws["logits"], ws["attn_lse"], o,
            _q_scale_dyn, kv_scale,
        )

    return o
scrolls · 307 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