submission 625902
brandonin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 461 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-625902?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:01cc174d510154c03b80925197d9b42065db74634642eb222b49f654acbe6d3c
license declaredunknown
license concludedunknown
authorsbrandonin
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
kv_raw, kv_scale = kv_data['mxfp4'] # fp4x2 tensor + fp8_e8m0 scaleshared-memory
extern __shared__ float smem[]; // [blockDim.x/32] floatsKernel source
submission.py461 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
# v12: Try mxfp4 KV path on top of v11.
# mxfp4 KV = 4-bit packed KV, half the bandwidth of fp8 KV (8-bit).
# For bandwidth-bound cases (bs=256, kv=8192), this is a 2x speedup opportunity.
# Strategy: probe get_mla_metadata_info_v1 with kv_dtype=fp4x2 at module load.
# If supported, use mxfp4 KV for all shapes; otherwise fall back to fp8 KV (v11).
import os
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ['CXX'] = 'clang++'
import sys
import torch
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
from torch.utils.cpp_extension import load_inline
NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_HEAD_DIM = 576
SM_SCALE = QK_HEAD_DIM ** -0.5
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
_FP8_FINFO = torch.finfo(FP8_DTYPE)
_FP8_MAX = _FP8_FINFO.max # 240.0 for e4m3fnuz, 448.0 for e4m3fn
# FP8 exponent bias: fnuz (max=240) has bias=8 → exp_offset=119; fn (max=448) has bias=7 → exp_offset=120
_FP8_EXP_OFFSET = 119 if _FP8_MAX < 300.0 else 120
# ─── HIP fused FP8 quantization kernel ──────────────────────────────────────
# Uses __hip_fp8_e4m3_fnuz type for correct FP8 conversion on gfx950.
# For small tensors (fits in one block): single-block kernel = 1 GPU dispatch.
# For large tensors: two-pass (amax reduction + quantize) = 2 GPU dispatches
# but still 1 Python-level dispatch.
HIP_SRC = r"""
#include <hip/hip_runtime.h>
#include <hip/hip_bfloat16.h>
#include <stdint.h>
#include <math.h>
// FP8 E4M3 conversion (handles both e4m3fnuz bias=8 and e4m3fn bias=7).
// exp_offset = 127 - fp8_bias = 119 for fnuz (bias=8), 120 for fn (bias=7).
// fp8_max = 240 for fnuz, 448 for fn.
// nan_bits = 0x80 for fnuz (no -0), 0x7F/0xFF for fn.
__device__ __forceinline__ uint8_t float_to_fp8_e4m3(float x, float fp8_max, int exp_offset) {
if (__isnanf(x)) return 0x80u; // treat as NaN regardless of variant
uint32_t bits;
__builtin_memcpy(&bits, &x, 4);
uint32_t sign = (bits >> 31) & 1u;
uint32_t exp32 = (bits >> 23) & 0xFFu;
uint32_t man32 = bits & 0x7FFFFFu;
// Zero (includes -0 which maps to +0 in fnuz; fn keeps -0 but that's rare)
if (exp32 == 0 && man32 == 0) return 0u;
// For e4m3fn (exp_offset=120, gfx950): 0x7F = NaN, max representable = 0x7E = 448.0
// For e4m3fnuz (exp_offset=119, gfx942): 0x7F = 240.0 = max, 0x80 = NaN
const uint8_t max_bits = (exp_offset == 120) ? (uint8_t)0x7Eu : (uint8_t)0x7Fu;
float ax = fabsf(x);
if (ax >= fp8_max) return (uint8_t)((sign << 7) | max_bits); // saturate
// fp8_biased_exp = (fp32_unbiased_exp) + fp8_bias
// = (exp32 - 127) + (127 - exp_offset)
// = exp32 - exp_offset
int fp8_exp = (int)exp32 - exp_offset;
if (fp8_exp <= 0) {
// Subnormal in fp8
int shift = 1 - fp8_exp;
if (shift > 4) return 0u; // underflow to zero
uint32_t man_full = man32 | 0x800000u;
uint32_t fp8_man = man_full >> (20 + shift);
uint32_t round_bit = (man_full >> (19 + shift)) & 1u;
if (round_bit) fp8_man++;
if (fp8_man > 7u) fp8_man = 7u;
return (uint8_t)((sign << 7) | fp8_man);
}
// Note: fp8_exp=15 is valid for both formats (covers values up to max).
// No early-exit here; saturation guard above handles values >= fp8_max.
// Normal: round mantissa 23→3 bits, round-to-nearest-even
uint32_t fp8_man = man32 >> 20;
uint32_t round_bit = (man32 >> 19) & 1u;
uint32_t sticky = man32 & ((1u << 19) - 1u);
if (round_bit && (sticky || (fp8_man & 1u))) {
fp8_man++;
if (fp8_man > 7u) {
fp8_man = 0u;
fp8_exp++;
}
}
uint32_t result = ((uint32_t)fp8_exp << 3) | fp8_man;
if (result > (uint32_t)max_bits) return (uint8_t)((sign << 7) | max_bits);
return (uint8_t)((sign << 7) | result);
}
// ─── Single-block kernel: handles numel ≤ blockDim.x * ITERS_PER_THREAD ───
// Phase 1: find global amax via shared memory reduction
// Phase 2: scale + cast to FP8
// → 1 kernel dispatch for small q tensors (bs=4: 36864 elements)
__global__ void fp8_quant_small_kernel(
const hip_bfloat16* __restrict__ src,
uint8_t* __restrict__ dst,
float* __restrict__ out_scale,
int numel,
float fp8_max,
int exp_offset
) {
extern __shared__ float smem[]; // [blockDim.x/32] floats
// Phase 1: thread-local max
float local_max = 0.0f;
for (int i = threadIdx.x; i < numel; i += blockDim.x) {
float v = fabsf((float)src[i]);
if (v > local_max) local_max = v;
}
// Warp reduce
for (int s = 16; s > 0; s >>= 1)
local_max = fmaxf(local_max, __shfl_down(local_max, s));
if (threadIdx.x % 32 == 0)
smem[threadIdx.x / 32] = local_max;
__syncthreads();
// Block reduce
if (threadIdx.x < 32) {
int nwarps = blockDim.x / 32;
float v = (threadIdx.x < nwarps) ? smem[threadIdx.x] : 0.0f;
for (int s = 16; s > 0; s >>= 1)
v = fmaxf(v, __shfl_down(v, s));
if (threadIdx.x == 0) {
float amax = fmaxf(v, 1e-12f);
smem[0] = amax;
out_scale[0] = amax / fp8_max;
}
}
__syncthreads();
// Phase 2: quantize
float inv_scale = fp8_max / smem[0];
for (int i = threadIdx.x; i < numel; i += blockDim.x) {
float v = (float)src[i] * inv_scale;
dst[i] = float_to_fp8_e4m3(v, fp8_max, exp_offset);
}
}
// ─── Two-pass approach for large tensors ───────────────────────────────────
// Pass 1: each block computes its block max, atomicMax to global_amax_bits[0]
__global__ void fp8_quant_amax_kernel(
const hip_bfloat16* __restrict__ src,
unsigned int* __restrict__ global_amax_bits, // must be initialized to 0
int numel
) {
extern __shared__ float smem[];
float local_max = 0.0f;
int stride = gridDim.x * blockDim.x;
for (int i = blockIdx.x * blockDim.x + threadIdx.x; i < numel; i += stride) {
float v = fabsf((float)src[i]);
if (v > local_max) local_max = v;
}
for (int s = 16; s > 0; s >>= 1)
local_max = fmaxf(local_max, __shfl_down(local_max, s));
if (threadIdx.x % 32 == 0) smem[threadIdx.x / 32] = local_max;
__syncthreads();
if (threadIdx.x < 32) {
int nwarps = blockDim.x / 32;
float v = (threadIdx.x < nwarps) ? smem[threadIdx.x] : 0.0f;
for (int s = 16; s > 0; s >>= 1)
v = fmaxf(v, __shfl_down(v, s));
if (threadIdx.x == 0) {
unsigned int ibits;
__builtin_memcpy(&ibits, &v, 4);
atomicMax(global_amax_bits, ibits);
}
}
}
// Pass 2: read global amax, write scale, quantize all elements
__global__ void fp8_quant_apply_kernel(
const hip_bfloat16* __restrict__ src,
uint8_t* __restrict__ dst,
float* __restrict__ out_scale,
unsigned int* __restrict__ global_amax_bits,
int numel,
float fp8_max,
int exp_offset
) {
unsigned int ibits = global_amax_bits[0];
float amax;
__builtin_memcpy(&amax, &ibits, 4);
amax = fmaxf(amax, 1e-12f);
float scale = amax / fp8_max;
if (blockIdx.x == 0 && threadIdx.x == 0) {
out_scale[0] = scale;
global_amax_bits[0] = 0u; // reset for next call
}
float inv_scale = 1.0f / scale;
int stride = gridDim.x * blockDim.x;
for (int i = blockIdx.x * blockDim.x + threadIdx.x; i < numel; i += stride) {
float v = (float)src[i] * inv_scale;
dst[i] = float_to_fp8_e4m3(v, fp8_max, exp_offset);
}
}
// C++ dispatch function: chooses single-block or two-pass
// _NUM_CU_CONST matches MI355X (304 compute units)
static const int _NUM_CU_CONST = 304;
void fused_fp8_quant(
torch::Tensor src, // [numel] bfloat16 (contiguous view)
torch::Tensor dst, // [numel] uint8 (pre-allocated)
torch::Tensor out_scale, // [1] float32 (pre-allocated)
torch::Tensor scratch, // [1] uint32 scratch for global amax (pre-allocated, init to 0)
double fp8_max_val,
int64_t exp_offset // 119 for e4m3fnuz (bias=8), 120 for e4m3fn (bias=7)
) {
int numel = src.numel();
float fp8_max = (float)fp8_max_val;
int exp_off = (int)exp_offset;
const int threads = 512;
if (numel <= threads * 256) {
// Single-block kernel: 1 GPU dispatch
int smem = (threads / 32) * sizeof(float);
fp8_quant_small_kernel<<<1, threads, smem>>>(
(const hip_bfloat16*)src.data_ptr(),
(uint8_t*)dst.data_ptr(),
out_scale.data_ptr<float>(),
numel,
fp8_max,
exp_off
);
} else {
// Two-pass: amax reduction + quantize (2 GPU dispatches, 1 Python call)
int blocks = min((numel + threads - 1) / threads, _NUM_CU_CONST * 4);
int smem = (threads / 32) * sizeof(float);
unsigned int* amax_ptr = (unsigned int*)scratch.data_ptr<int>();
fp8_quant_amax_kernel<<<blocks, threads, smem>>>(
(const hip_bfloat16*)src.data_ptr(),
amax_ptr,
numel
);
fp8_quant_apply_kernel<<<blocks, threads, 0>>>(
(const hip_bfloat16*)src.data_ptr(),
(uint8_t*)dst.data_ptr(),
out_scale.data_ptr<float>(),
amax_ptr,
numel,
fp8_max,
exp_off
);
}
}
"""
CPP_SRC = """
#include <torch/extension.h>
void fused_fp8_quant(
torch::Tensor src,
torch::Tensor dst,
torch::Tensor out_scale,
torch::Tensor scratch,
double fp8_max_val,
int64_t exp_offset
);
"""
_quant_module = None
_fused_fp8_quant = None
try:
_quant_module = load_inline(
name='mixed_mla_fp8_quant_v12',
cpp_sources=[CPP_SRC],
cuda_sources=[HIP_SRC],
functions=['fused_fp8_quant'],
verbose=False,
extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3",
"-mcumode"],
)
_fused_fp8_quant = _quant_module.fused_fp8_quant
print("[v10] FP8 quant kernel loaded OK", file=sys.stderr)
except Exception as e:
print(f"[v10] FP8 quant kernel compile FAILED: {e}", file=sys.stderr)
# ─── Shape metadata cache ────────────────────────────────────────────────────
_NUM_KV_SPLITS = 32
# mxfp4 KV not supported: mla_decode_stage1_asm_fwd requires head_size == KV.size(3)
# but fp4x2 KV has KV.size(3)=288 while head_size=576. Always use fp8 KV.
_USE_MXFP4_KV = False
_KV_DTYPE = FP8_DTYPE
_KV_HEAD_DIM = QK_HEAD_DIM # 576 for fp8
def _fallback_quantize_fp8(tensor):
"""Fallback to PyTorch ops if HIP kernel failed to compile."""
amax = tensor.abs().amax().clamp(min=1e-12)
scale = amax / _FP8_FINFO.max
fp8_tensor = (tensor / scale).clamp(_FP8_FINFO.min, _FP8_FINFO.max).to(FP8_DTYPE)
return fp8_tensor, scale.to(torch.float32).reshape(1)
# shape_key -> ShapeCache
_shape_cache = {}
class _ShapeCache:
"""Per-shape pre-allocated buffers + metadata."""
__slots__ = [
'q_fp8_buf', 'q_scale_buf', 'scratch_buf',
'out', 'meta', 'kv_indices', 'kv_last_page_len',
'qo_indptr_cached', 'kv_indptr_cached',
'q_numel',
]
def __init__(self, batch_size, kv_seq_len, q_seq_len, qo_indptr, kv_indptr):
total_q = batch_size * q_seq_len
q_numel = total_q * NUM_HEADS * QK_HEAD_DIM
self.q_fp8_buf = torch.empty(q_numel, dtype=torch.uint8, device='cuda')
self.q_scale_buf = torch.zeros(1, dtype=torch.float32, device='cuda')
self.scratch_buf = torch.zeros(1, dtype=torch.int32, device='cuda')
self.q_numel = q_numel
self.out = torch.empty(
(total_q, NUM_HEADS, KV_LORA_RANK),
dtype=torch.bfloat16, device='cuda',
)
total_kv = int(kv_indptr[-1].item())
self.kv_indices = torch.arange(total_kv, dtype=torch.int32, device='cuda')
self.kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
self.qo_indptr_cached = qo_indptr.clone()
self.kv_indptr_cached = kv_indptr.clone()
info = get_mla_metadata_info_v1(
batch_size, q_seq_len, NUM_HEADS, FP8_DTYPE, _KV_DTYPE,
is_sparse=False, fast_mode=False,
num_kv_splits=_NUM_KV_SPLITS, intra_batch_mode=True,
)
work = [torch.empty(s, dtype=t, device='cuda') for s, t in info]
(work_metadata, work_indptr, work_info_set,
reduce_indptr, reduce_final_map, reduce_partial_map) = work
get_mla_metadata_v1(
qo_indptr, kv_indptr, self.kv_last_page_len,
NUM_HEADS // NUM_KV_HEADS,
NUM_KV_HEADS,
True,
work_metadata, work_info_set, work_indptr,
reduce_indptr, reduce_final_map, reduce_partial_map,
page_size=PAGE_SIZE,
kv_granularity=max(PAGE_SIZE, 16),
max_seqlen_qo=q_seq_len,
uni_seqlen_qo=q_seq_len,
fast_mode=False,
max_split_per_batch=_NUM_KV_SPLITS,
intra_batch_mode=True,
dtype_q=FP8_DTYPE,
dtype_kv=_KV_DTYPE,
)
self.meta = {
'work_meta_data': work_metadata,
'work_indptr': work_indptr,
'work_info_set': work_info_set,
'reduce_indptr': reduce_indptr,
'reduce_final_map': reduce_final_map,
'reduce_partial_map': reduce_partial_map,
}
print(f"[v12] ShapeCache ({batch_size},{kv_seq_len},{q_seq_len}) "
f"kv_dtype={'fp4x2' if _USE_MXFP4_KV else 'fp8'} splits={_NUM_KV_SPLITS}",
file=sys.stderr)
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config['batch_size']
kv_seq_len = config['kv_seq_len']
q_seq_len = config['q_seq_len']
shape_key = (batch_size, kv_seq_len, q_seq_len)
# Build/reuse per-shape cache
if shape_key not in _shape_cache:
_shape_cache[shape_key] = _ShapeCache(
batch_size, kv_seq_len, q_seq_len, qo_indptr, kv_indptr
)
sc = _shape_cache[shape_key]
# ── Q FP8 quantization ────────────────────────────────────────────────────
q_flat = q.contiguous().view(-1)
if _fused_fp8_quant is not None:
_fused_fp8_quant(q_flat, sc.q_fp8_buf, sc.q_scale_buf, sc.scratch_buf,
_FP8_MAX, _FP8_EXP_OFFSET)
q_fp8 = sc.q_fp8_buf.view(FP8_DTYPE)
q_scale = sc.q_scale_buf
else:
q_fp8_raw, q_scale = _fallback_quantize_fp8(q)
q_fp8 = q_fp8_raw.view(-1)
# ── KV selection ─────────────────────────────────────────────────────────
if _USE_MXFP4_KV:
kv_raw, kv_scale = kv_data['mxfp4'] # fp4x2 tensor + fp8_e8m0 scale
total_kv = kv_raw.shape[0]
kv_4d = kv_raw.view(total_kv, PAGE_SIZE, NUM_KV_HEADS, _KV_HEAD_DIM)
else:
kv_raw, kv_scale = kv_data['fp8']
total_kv = kv_raw.shape[0]
kv_4d = kv_raw.view(total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)
# ── MLA decode ────────────────────────────────────────────────────────────
mla_decode_fwd(
q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM),
kv_4d,
sc.out,
sc.qo_indptr_cached,
sc.kv_indptr_cached,
sc.kv_indices,
sc.kv_last_page_len,
q_seq_len,
page_size=PAGE_SIZE,
nhead_kv=NUM_KV_HEADS,
sm_scale=SM_SCALE,
logit_cap=0.0,
num_kv_splits=_NUM_KV_SPLITS,
q_scale=q_scale,
kv_scale=kv_scale,
intra_batch_mode=True,
**sc.meta,
)
return sc.out
scrolls · 461 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