submission 697838
Zhenyu2Liang · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 291 lines, June 9 Researcher Reciprocity License v1.0.
submission_v78.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-697838?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:363387f7764a8d9c4a45a993837fd01396d8ba21c2a392b1fd13f3c915708e74
license declaredunknown
license concludedunknown
authorsZhenyu2Liang
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
vector-width = uint4
const uint4* vec_input = reinterpret_cast<const uint4*>(input_bytes);Kernel source
submission_v78.py291 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
MLA V78 - True Unaligned Native Vectorization & Occupancy Tuning
- Unlocks 100% native v_load_b128 load logic by leveraging PyTorch strictly aligned 256-bye allocator boundary.
- Cranks Wavefront thread scheduling block size from 64 to 256, maximizing CU hardware cache usage.
- Enhances large contextual batching mappings in SPLIT_MAP.
"""
import os
import torch
from task import input_t, output_t
import aiter
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
os.environ['PYTORCH_ROCM_ARCH'] = 'gfx950'
os.environ['CXX'] = 'clang++'
CPP_SRC = """
#include <torch/extension.h>
#include <c10/core/DispatchKeySet.h>
#include <ATen/core/dispatch/Dispatcher.h>
#include <tuple>
#include <string>
void fast_fp8_cast_and_indptr_inline(torch::Tensor q, torch::Tensor q_fp8, float scale,
torch::Tensor kv_indptr, torch::Tensor kv_last_page_len);
"""
CUDA_SRC = """
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <unordered_map>
#include <string>
#include <c10/core/DispatchKeySet.h>
#include <ATen/core/dispatch/Dispatcher.h>
#include <cstdint>
union FloatInt { float f; uint32_t i; };
__device__ __forceinline__ uint8_t cast_to_fp8_e4m3fn(float val) {
if (val == 0.0f) return 0x00;
if (val != val) return 0x7F;
FloatInt val_u; val_u.f = val; uint32_t bits = val_u.i;
uint32_t sign = (bits >> 24) & 0x80;
int32_t exp = ((bits >> 23) & 0xFF) - 127;
uint32_t mantissa = bits & 0x7FFFFF;
if (exp == 128) return sign | 0x7F;
int32_t target_exp = exp + 7;
if (target_exp <= 0) {
int32_t shift = 1 - target_exp;
if (shift > 4) return sign;
uint32_t m = (1 << 23) | mantissa;
uint32_t mask = (1 << (shift + 20)) - 1;
uint32_t half = 1 << (shift + 19);
uint32_t remainder = m & mask;
m >>= (shift + 20);
if (remainder > half || (remainder == half && (m & 1))) { m += 1; }
if (m == 8) { return sign | 0x08; }
return sign | m;
}
if (target_exp >= 16) { return sign | 0x7E; }
uint32_t mask = (1 << 20) - 1;
uint32_t half = 1 << 19;
uint32_t remainder = mantissa & mask;
uint32_t m = mantissa >> 20;
if (remainder > half || (remainder == half && (m & 1))) {
m += 1; if (m == 8) { target_exp++; m = 0; }
}
if (target_exp >= 16) { return sign | 0x7E; }
if (target_exp == 15 && m == 7) { m = 6; }
return sign | (target_exp << 3) | m;
}
// ---------------------------------------------------------
// V78 Ultimate Aligned Vectorization & 256-Thread Cast Core
// ---------------------------------------------------------
__global__ __launch_bounds__(256)
void fp8_cast_unaligned_128(
const uint8_t* __restrict__ input_bytes,
uint2* __restrict__ output_vec,
float scale, int total_vecs,
const int32_t* __restrict__ kv_indptr,
int32_t* __restrict__ kv_last_page_len, int batch_size,
const uint16_t* __restrict__ input_scalar,
uint8_t* __restrict__ output_scalar,
int total_elements) {
int idx = blockIdx.x * blockDim.x + threadIdx.x;
if (idx < total_vecs) {
// V78: Removed __attribute__((packed)) since PyTorch pointers are natively strongly aligned
// Generates native 128-bit vector v_load_b128!
const uint4* vec_input = reinterpret_cast<const uint4*>(input_bytes);
uint4 val_packed = vec_input[idx];
uint16_t in_arr[8];
in_arr[0] = val_packed.x & 0xFFFF; in_arr[1] = val_packed.x >> 16;
in_arr[2] = val_packed.y & 0xFFFF; in_arr[3] = val_packed.y >> 16;
in_arr[4] = val_packed.z & 0xFFFF; in_arr[5] = val_packed.z >> 16;
in_arr[6] = val_packed.w & 0xFFFF; in_arr[7] = val_packed.w >> 16;
uint32_t out_arr_32[2] = {0, 0};
uint8_t* out_ptr = (uint8_t*)out_arr_32;
#pragma unroll 8
for(int i=0; i<8; i++) {
FloatInt val_u; val_u.i = ((uint32_t)in_arr[i]) << 16;
float val = val_u.f * scale;
val = fmaxf(val, -448.0f); val = fminf(val, 448.0f);
out_ptr[i] = cast_to_fp8_e4m3fn(val);
}
uint2 out_val; out_val.x = out_arr_32[0]; out_val.y = out_arr_32[1];
// Emits single ultra-fast natively aligned 64-bit store instruction without decoder pressure.
output_vec[idx] = out_val;
}
// Tail coverage safely resolved dynamically
if (idx == 0) {
int tail_start = total_vecs * 8;
for (int i = tail_start; i < total_elements; i++) {
FloatInt val_u; val_u.i = ((uint32_t)input_scalar[i]) << 16;
float val = val_u.f * scale;
val = fmaxf(val, -448.0f); val = fminf(val, 448.0f);
output_scalar[i] = cast_to_fp8_e4m3fn(val);
}
}
// Non-divergent fast pointer alignment matrix indexing
if (idx < batch_size) {
kv_last_page_len[idx] = kv_indptr[idx + 1] - kv_indptr[idx];
}
}
// ---------------------------------------------------------
// Global Bound Vectorization Override Launcher
// ---------------------------------------------------------
void fast_fp8_cast_and_indptr_inline(torch::Tensor q, torch::Tensor q_fp8, float scale,
torch::Tensor kv_indptr, torch::Tensor kv_last_page_len) {
int total_elements = q.numel();
int batch_size = kv_last_page_len.size(0);
int total_vecs = total_elements / 8;
int threads_needed = (total_vecs > batch_size) ? total_vecs : batch_size;
int threads = 256;
int blocks = (threads_needed + threads - 1) / threads;
if (blocks == 0 && batch_size > 0) { blocks = 1; }
if (blocks > 0) {
fp8_cast_unaligned_128<<<blocks, threads>>>(
reinterpret_cast<const uint8_t*>(q.data_ptr()),
reinterpret_cast<uint2*>(q_fp8.data_ptr()),
scale, total_vecs,
kv_indptr.data_ptr<int32_t>(), kv_last_page_len.data_ptr<int32_t>(), batch_size,
reinterpret_cast<const uint16_t*>(q.data_ptr()),
reinterpret_cast<uint8_t*>(q_fp8.data_ptr()), total_elements
);
}
}
"""
_module_cache = None
def get_module():
global _module_cache
if _module_cache is None:
from torch.utils.cpp_extension import load_inline
_module_cache = load_inline(
name="v78_mla_native_turbo", cpp_sources=[CPP_SRC], cuda_sources=[CUDA_SRC],
functions=["fast_fp8_cast_and_indptr_inline"], with_cuda=True,
extra_cflags=["-O3"], extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3"],
verbose=False, build_directory="/tmp"
)
return _module_cache
FP8_DTYPE = aiter_dtypes.fp8
PAGE_SIZE = 1
_workspace_cache = {}
_out_cache = {}
_q_input_cache = {}
_kv_idx_cache = {}
_kv_last_page_cache = {}
_target_q_scale = None
_METADATA_DONE = set()
_SPLIT_MAP = { 4: 76, 32: 19, 64: 9, 256: 2 }
_FAST_MODULE = None
def custom_kernel(data: input_t) -> output_t:
global _target_q_scale, _FAST_MODULE
if _FAST_MODULE is None:
_FAST_MODULE = get_module().fast_fp8_cast_and_indptr_inline
_target_q_scale = torch.empty((1,), dtype=torch.float32, device="cuda").fill_(1.0 / 448.0)
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = int(config["batch_size"])
q_seq_len = int(config["q_seq_len"])
kv_input, kv_scale = kv_data["fp8"]
out_key = (batch_size, q_seq_len)
if out_key not in _q_input_cache:
buf = torch.empty_like(q, dtype=FP8_DTYPE)
_q_input_cache[out_key] = buf
_kv_last_page_cache[batch_size] = torch.empty(batch_size, dtype=torch.int32, device="cuda")
q_input = _q_input_cache[out_key]
kv_last_page_len = _kv_last_page_cache[batch_size]
_FAST_MODULE(q, q_input, 448.0, kv_indptr, kv_last_page_len)
kv_seq_len = int(config["kv_seq_len"])
total_kv_len = batch_size * kv_seq_len
if total_kv_len not in _kv_idx_cache:
_kv_idx_cache[total_kv_len] = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
kv_indices = _kv_idx_cache[total_kv_len]
nq = int(config["num_heads"])
dq = int(config["qk_head_dim"])
nkv = int(config["num_kv_heads"])
kv_buffer_4d = kv_input.view(kv_input.shape[0], PAGE_SIZE, nkv, kv_input.shape[-1])
is_fast = (batch_size >= 32)
gran = 32 if is_fast else 16
num_kv_splits = _SPLIT_MAP.get(batch_size, 8)
ws_key = (batch_size, q_seq_len, num_kv_splits, is_fast)
if ws_key not in _workspace_cache:
info = get_mla_metadata_info_v1(
batch_size, q_seq_len, nq, FP8_DTYPE, kv_input.dtype,
is_sparse=False, fast_mode=is_fast,
num_kv_splits=num_kv_splits, intra_batch_mode=True,
)
work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
_workspace_cache[ws_key] = work
work_metadata, work_indptr, work_info_set, reduce_indptr, reduce_final_map, reduce_partial_map = _workspace_cache[ws_key]
bind_key = (batch_size, q_seq_len, kv_seq_len, num_kv_splits, is_fast, gran)
if bind_key not in _METADATA_DONE:
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_last_page_len, nq // nkv, nkv, True,
work_metadata, work_info_set, work_indptr, reduce_indptr, reduce_final_map, reduce_partial_map,
page_size=PAGE_SIZE, kv_granularity=gran, max_seqlen_qo=q_seq_len, uni_seqlen_qo=q_seq_len,
fast_mode=is_fast, max_split_per_batch=num_kv_splits, intra_batch_mode=True,
dtype_q=FP8_DTYPE, dtype_kv=kv_input.dtype,
)
_METADATA_DONE.add(bind_key)
if out_key not in _out_cache:
dv = int(config["v_head_dim"])
_out_cache[out_key] = torch.empty((q.shape[0], nq, dv), dtype=torch.bfloat16, device="cuda")
o = _out_cache[out_key]
mla_decode_fwd(
q_input.view(-1, nq, dq),
kv_buffer_4d,
o,
qo_indptr,
kv_indptr,
kv_indices,
kv_last_page_len,
q_seq_len,
page_size=PAGE_SIZE,
nhead_kv=int(config["num_kv_heads"]),
sm_scale=float(config["sm_scale"]),
logit_cap=0.0,
num_kv_splits=num_kv_splits,
q_scale=_target_q_scale,
kv_scale=kv_scale,
intra_batch_mode=True,
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,
)
return o
scrolls · 291 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