Skip to content
KernelIndex
Search⌘K

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
AMD Instinct MI355X
86.7µs
#412 of 766
2026-03-25

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.

fp4kv_raw, kv_scale = kv_data['mxfp4'] # fp4x2 tensor + fp8_e8m0 scale
shared-memoryextern __shared__ float smem[]; // [blockDim.x/32] floats

Kernel 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