Skip to content
KernelIndex
Search⌘K

submission 673989

darkness098600 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 439 lines, June 9 Researcher Reciprocity License v1.0.

submission_v36.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-673989?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
93.0µs
#432 of 766
2026-03-30

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:96aa6eb891b564165d167d8ec14e8f0e21323581fe56fdad169b767d72ee3008
license declaredunknown
license concludedunknown
authorsdarkness098600
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

online-softmax__shared__ float running_max[kHeadsPerBlock];
shared-memory__shared__ float q_shared[kHeadsPerBlock * kQKDim];
split-kvoid mla_splitk_stage1_fp8(

Kernel source

submission_v36.py439 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X

"""
Mixed-MLA v36 - intra_batch_mode=False 测试

v29: intra_batch_mode=True → benchmark 90.7µs, leaderboard 94.2µs
测试 intra_batch_mode=False 是否能改善性能
"""

import math
import os

os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("CXX", "clang++")

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t


NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
SM_SCALE = 1.0 / math.sqrt(QK_HEAD_DIM)
PAGE_SIZE = 1
NUM_KV_SPLITS = 24
NATIVE_BATCH_LIMIT = 4

_RUNTIME = None
_META_CACHE = {}
_NATIVE_MOD = None

_NATIVE_CPP = r"""
void mla_splitk_stage1_fp8(
    torch::Tensor q,
    torch::Tensor kv,
    torch::Tensor kv_indptr,
    torch::Tensor kv_scale,
    torch::Tensor partial_m,
    torch::Tensor partial_l,
    torch::Tensor partial_acc,
    int num_splits
);
"""

_NATIVE_HIP = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <hip/amd_detail/amd_hip_fp8.h>
#include <cmath>

constexpr int kNumHeads = 16;
constexpr int kHeadsPerBlock = 8;
constexpr int kQKDim = 576;
constexpr int kVDim = 512;
constexpr int kThreads = 512;
constexpr int kHeadWidth = 64;
constexpr int kValsPerLane = kVDim / kHeadWidth;
constexpr float kSmScale = 1.0f / 24.0f;

__device__ __forceinline__ float decode_e4m3_ocp(__hip_fp8_storage_t x) {
    union {
        unsigned int u32;
        unsigned char u8[4];
    } bits;
    bits.u32 = 0;
    bits.u8[0] = x;
    return __builtin_amdgcn_cvt_f32_fp8(bits.u32, 0);
}

__global__ void mla_splitk_stage1_mqa_kernel(
    const __hip_bfloat16* q,
    const __hip_fp8_storage_t* kv,
    const int32_t* kv_indptr,
    const float* kv_scale,
    float* partial_m,
    float* partial_l,
    float* partial_acc,
    int num_splits
) {
    const int batch_id = static_cast<int>(blockIdx.x);
    const int split_id = static_cast<int>(blockIdx.y);
    const int head_group = static_cast<int>(blockIdx.z);
    const int tid = static_cast<int>(threadIdx.x);
    const int local_head_id = tid / kHeadWidth;
    const int lane = tid % kHeadWidth;
    const int head_id = head_group * kHeadsPerBlock + local_head_id;

    __shared__ float q_shared[kHeadsPerBlock * kQKDim];
    __shared__ float kv_shared[kQKDim];
    __shared__ float running_max[kHeadsPerBlock];
    __shared__ float running_sum[kHeadsPerBlock];
    __shared__ float coeff_old[kHeadsPerBlock];
    __shared__ float coeff_new[kHeadsPerBlock];

    const int kv_start_total = kv_indptr[batch_id];
    const int kv_end_total = kv_indptr[batch_id + 1];
    const int kv_len = kv_end_total - kv_start_total;
    const int split_start = kv_start_total + (kv_len * split_id) / num_splits;
    const int split_end = kv_start_total + (kv_len * (split_id + 1)) / num_splits;

    const int q_group_base = (batch_id * kNumHeads + head_group * kHeadsPerBlock) * kQKDim;
    for (int idx = tid; idx < kHeadsPerBlock * kQKDim; idx += kThreads) {
        q_shared[idx] = __bfloat162float(q[q_group_base + idx]);
    }
    if (tid < kHeadsPerBlock) {
        running_max[tid] = -INFINITY;
        running_sum[tid] = 0.0f;
        coeff_old[tid] = 0.0f;
        coeff_new[tid] = 0.0f;
    }
    __syncthreads();

    float acc[kValsPerLane];
    #pragma unroll
    for (int i = 0; i < kValsPerLane; ++i) {
        acc[i] = 0.0f;
    }

    if (split_start >= split_end) {
        if (lane == 0) {
            partial_m[(batch_id * kNumHeads + head_id) * num_splits + split_id] = -INFINITY;
            partial_l[(batch_id * kNumHeads + head_id) * num_splits + split_id] = 0.0f;
        }
        const int acc_base = ((batch_id * kNumHeads + head_id) * num_splits + split_id) * kVDim;
        #pragma unroll
        for (int i = 0; i < kValsPerLane; ++i) {
            partial_acc[acc_base + lane + i * kHeadWidth] = 0.0f;
        }
        return;
    }

    const float scale = kv_scale[0];
    for (int kv_idx = split_start; kv_idx < split_end; ++kv_idx) {
        const __hip_fp8_storage_t* kv_ptr = kv + kv_idx * kQKDim;
        for (int idx = tid; idx < kQKDim; idx += kThreads) {
            kv_shared[idx] = decode_e4m3_ocp(kv_ptr[idx]) * scale;
        }
        __syncthreads();

        float partial = 0.0f;
        const int q_base = local_head_id * kQKDim;
        for (int d = lane; d < kQKDim; d += kHeadWidth) {
            partial += q_shared[q_base + d] * kv_shared[d];
        }
        #pragma unroll
        for (int offset = kHeadWidth / 2; offset > 0; offset >>= 1) {
            partial += __shfl_down(partial, offset, kHeadWidth);
        }

        if (lane == 0) {
            const float score = partial * kSmScale;
            const float prev_max = running_max[head_id];
            const float prev_sum = running_sum[head_id];
            const float next_max = fmaxf(prev_max, score);
            const float prev_scale = (prev_sum == 0.0f) ? 0.0f : expf(prev_max - next_max);
            const float next_term = expf(score - next_max);
            running_max[head_id] = next_max;
            running_sum[head_id] = prev_sum * prev_scale + next_term;
            coeff_old[head_id] = prev_scale;
            coeff_new[head_id] = next_term;
        }
        __syncthreads();

        const float alpha = coeff_old[head_id];
        const float beta = coeff_new[head_id];
        #pragma unroll
        for (int i = 0; i < kValsPerLane; ++i) {
            const int v_idx = lane + i * kHeadWidth;
            acc[i] = acc[i] * alpha + beta * kv_shared[v_idx];
        }
        __syncthreads();
    }

    if (lane == 0) {
        partial_m[(batch_id * kNumHeads + head_id) * num_splits + split_id] = running_max[head_id];
        partial_l[(batch_id * kNumHeads + head_id) * num_splits + split_id] = running_sum[head_id];
    }
    const int acc_base = ((batch_id * kNumHeads + head_id) * num_splits + split_id) * kVDim;
    #pragma unroll
    for (int i = 0; i < kValsPerLane; ++i) {
        partial_acc[acc_base + lane + i * kHeadWidth] = acc[i];
    }
}

void mla_splitk_stage1_fp8(
    torch::Tensor q,
    torch::Tensor kv,
    torch::Tensor kv_indptr,
    torch::Tensor kv_scale,
    torch::Tensor partial_m,
    torch::Tensor partial_l,
    torch::Tensor partial_acc,
    int num_splits
) {
    const int batch_size = static_cast<int>(q.size(0));
    hipLaunchKernelGGL(
        mla_splitk_stage1_mqa_kernel,
        dim3(batch_size, num_splits, kNumHeads / kHeadsPerBlock),
        dim3(kThreads),
        0,
        0,
        reinterpret_cast<const __hip_bfloat16*>(q.data_ptr()),
        reinterpret_cast<const __hip_fp8_storage_t*>(kv.data_ptr()),
        kv_indptr.data_ptr<int32_t>(),
        kv_scale.data_ptr<float>(),
        partial_m.data_ptr<float>(),
        partial_l.data_ptr<float>(),
        partial_acc.data_ptr<float>(),
        num_splits
    );
}
"""


def _get_native_mod():
    global _NATIVE_MOD
    if _NATIVE_MOD is None:
        _NATIVE_MOD = load_inline(
            name="mixed_mla_splitk_v36",
            cpp_sources=[_NATIVE_CPP],
            cuda_sources=[_NATIVE_HIP],
            functions=["mla_splitk_stage1_fp8"],
            extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20", "-O3"],
            verbose=False,
        )
    return _NATIVE_MOD


def _ensure_runtime():
    global _RUNTIME
    if _RUNTIME is not None:
        return _RUNTIME

    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

    fp8_info = torch.finfo(aiter_dtypes.fp8)
    _RUNTIME = (
        mla_decode_fwd,
        aiter_dtypes.fp8,
        fp8_info.min,
        fp8_info.max,
        get_mla_metadata_info_v1,
        get_mla_metadata_v1,
    )
    return _RUNTIME


def _quantize_fp8(q: torch.Tensor, fp8_dtype: torch.dtype, fp8_min: float, fp8_max: float):
    amax = q.abs().amax().clamp(min=1e-12)
    scale = (amax / fp8_max).to(torch.float32).view(1)
    q_fp8 = (q / scale).clamp(min=fp8_min, max=fp8_max).to(fp8_dtype)
    return q_fp8, scale


def _get_cached_meta(
    batch_size: int,
    kv_seq_len: int,
    q_dtype: torch.dtype,
    kv_dtype: torch.dtype,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    get_mla_metadata_info_v1,
    get_mla_metadata_v1,
):
    key = (batch_size, kv_seq_len, str(q_dtype), str(kv_dtype), NUM_KV_SPLITS, qo_indptr.device.index)
    cached = _META_CACHE.get(key)
    if cached is not None:
        return cached

    max_q_len = 1
    kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
    info = get_mla_metadata_info_v1(
        batch_size,
        max_q_len,
        NUM_HEADS,
        q_dtype,
        kv_dtype,
        is_sparse=False,
        fast_mode=False,
        num_kv_splits=NUM_KV_SPLITS,
        intra_batch_mode=False,
    )
    work = [torch.empty(shape, dtype=dtype, device="cuda") for shape, dtype 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,
        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=max_q_len,
        uni_seqlen_qo=max_q_len,
        fast_mode=False,
        max_split_per_batch=NUM_KV_SPLITS,
        intra_batch_mode=False,
        dtype_q=q_dtype,
        dtype_kv=kv_dtype,
    )

    cached = {
        "kv_indices": torch.arange(batch_size * kv_seq_len, dtype=torch.int32, device="cuda"),
        "kv_last_page_len": kv_last_page_len,
        "work": {
            "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,
        },
    }
    _META_CACHE[key] = cached
    return cached


def _aiter_path(
    q: torch.Tensor,
    kv_data: dict,
    qo_indptr: torch.Tensor,
    kv_indptr: torch.Tensor,
    config: dict,
) -> torch.Tensor:
    (
        mla_decode_fwd,
        fp8_dtype,
        fp8_min,
        fp8_max,
        get_mla_metadata_info_v1,
        get_mla_metadata_v1,
    ) = _ensure_runtime()

    q_input, q_scale = _quantize_fp8(q, fp8_dtype, fp8_min, fp8_max)
    kv_input, kv_scale = kv_data["fp8"]
    meta = _get_cached_meta(
        config["batch_size"],
        config["kv_seq_len"],
        q_input.dtype,
        kv_input.dtype,
        qo_indptr,
        kv_indptr,
        get_mla_metadata_info_v1,
        get_mla_metadata_v1,
    )

    out = torch.empty((q.shape[0], NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
    mla_decode_fwd(
        q_input,
        kv_input.view(kv_input.shape[0], PAGE_SIZE, NUM_KV_HEADS, kv_input.shape[-1]),
        out,
        qo_indptr,
        kv_indptr,
        meta["kv_indices"],
        meta["kv_last_page_len"],
        1,
        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=False,
        **meta["work"],
    )
    return out


def _native_num_splits(kv_seq_len: int) -> int:
    return 128 if kv_seq_len >= 8192 else 0


# v22: 关键修复 - batch=4 且 seq=8192 时不使用 native path
def _should_use_native(config: dict) -> bool:
    # batch=4, seq=8192 时 native path 反而慢,回退到 aiter
    if config["q_seq_len"] == 1 and config["batch_size"] <= NATIVE_BATCH_LIMIT and config["kv_seq_len"] >= 8192:
        # 唯一的问题 case: batch=4 且 seq=8192
        if config["batch_size"] == 4 and config["kv_seq_len"] == 8192:
            return False  # 使用 aiter
        return True
    return False


def _native_path(q: torch.Tensor, kv_data: dict, kv_indptr: torch.Tensor, config: dict) -> torch.Tensor:
    kv_fp8, kv_scale = kv_data["fp8"]
    batch_size = config["batch_size"]
    num_splits = _native_num_splits(config["kv_seq_len"])

    partial_m = torch.empty((batch_size, NUM_HEADS, num_splits), dtype=torch.float32, device=q.device)
    partial_l = torch.empty_like(partial_m)
    partial_acc = torch.empty((batch_size, NUM_HEADS, num_splits, V_HEAD_DIM), dtype=torch.float32, device=q.device)

    _get_native_mod().mla_splitk_stage1_fp8(
        q.contiguous(),
        kv_fp8.view(-1, QK_HEAD_DIM).contiguous(),
        kv_indptr.contiguous(),
        kv_scale.contiguous(),
        partial_m,
        partial_l,
        partial_acc,
        num_splits,
    )

    max_m = partial_m.max(dim=2, keepdim=True).values
    weights = torch.exp(partial_m - max_m)
    denom = (partial_l * weights).sum(dim=2, keepdim=True).clamp_min(1e-12)
    out = (partial_acc * weights.unsqueeze(-1)).sum(dim=2) / denom
    return out.to(torch.bfloat16)


def custom_kernel(data: input_t) -> output_t:
    q, kv_data, qo_indptr, kv_indptr, config = data
    if _should_use_native(config):
        return _native_path(q, kv_data, kv_indptr, config)
    return _aiter_path(q, kv_data, qo_indptr, kv_indptr, config)
scrolls · 439 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