Skip to content
KernelIndex
Search⌘K

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
AMD Instinct MI355X
58.1µs
#227 of 766
2026-04-02

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 = uint4const 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