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
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-k
void 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