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
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-kernel
kv<=1024: persistent mode (32-split + mla_reduce_v1) matching reference exactlyshared-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