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
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-kernel
NUM_KV_SPLITS = 32 # fixed — aiter persistent kernel needs enough work unitsshared-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