submission 646958
Maxwell Cipher · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 239 lines, June 9 Researcher Reciprocity License v1.0.
c_v17.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-646958?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:a642fbdf86dfed0ef093305a2727d9ec6c404660e12d15df40cd34a85d001fd6
license declaredunknown
license concludedunknown
authorsMaxwell Cipher
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
persistent-kernel
PERSISTENT_TOKEN_LIMIT = 65536Kernel source
c_v17.py239 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
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
try:
from aiter import per_tensor_quant_hip
HAS_QUANT_HIP = True
except ImportError:
HAS_QUANT_HIP = False
NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
PAGE_SIZE = 1
DEFAULT_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
FP8_DTYPE = aiter_dtypes.fp8
LOW_PAYLOAD_TOKEN_LIMIT = 65536
PERSISTENT_TOKEN_LIMIT = 65536
PERSISTENT_SPLITS = 32
def _quantize_query(q_view: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
if HAS_QUANT_HIP:
q_fp8, q_scale = per_tensor_quant_hip(q_view, quant_dtype=FP8_DTYPE)
return q_fp8, q_scale.reshape(1)
finfo = torch.finfo(FP8_DTYPE)
amax = q_view.abs().amax().clamp(min=1e-12)
q_scale = (amax / finfo.max).to(torch.float32).reshape(1)
q_fp8 = (q_view / q_scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
return q_fp8, q_scale
def _prefer_low_overhead_path(total_kv: int) -> bool:
return total_kv <= LOW_PAYLOAD_TOKEN_LIMIT
def _prefer_persistent_fp8_path(total_kv: int) -> bool:
return total_kv > PERSISTENT_TOKEN_LIMIT
def _run_nonpersistent_path(
q_view: torch.Tensor,
kv_data: dict,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
batch_size: int,
kv_seq_len: int,
sm_scale: float,
) -> torch.Tensor:
total_q = q_view.shape[0]
total_kv = batch_size * kv_seq_len
device = q_view.device
page_ids = torch.arange(total_kv, dtype=torch.int32, device=device)
last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)
output = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device)
if _prefer_low_overhead_path(total_kv):
kv_tensor = kv_data["bf16"]
q_input = q_view
q_scale = None
kv_scale = None
else:
kv_tensor, kv_scale = kv_data["fp8"]
q_input, q_scale = _quantize_query(q_view)
kv_pages = kv_tensor.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_tensor.shape[-1])
mla_decode_fwd(
q_input,
kv_pages,
output,
qo_indptr,
kv_indptr,
page_ids,
last_page_len,
1,
page_size=PAGE_SIZE,
nhead_kv=NUM_KV_HEADS,
sm_scale=sm_scale,
q_scale=q_scale,
kv_scale=kv_scale,
intra_batch_mode=False,
)
return output
def _build_persistent_metadata(
batch_size: int,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
q_dtype: torch.dtype,
kv_dtype: torch.dtype,
kv_last_page_len: torch.Tensor,
):
info = get_mla_metadata_info_v1(
batch_size,
1,
NUM_HEADS,
q_dtype,
kv_dtype,
is_sparse=False,
fast_mode=False,
num_kv_splits=PERSISTENT_SPLITS,
intra_batch_mode=True,
)
work_meta_data, work_indptr, work_info_set, reduce_indptr, reduce_final_map, reduce_partial_map = [
torch.empty(shape, dtype=dtype, device=qo_indptr.device) for shape, dtype in info
]
get_mla_metadata_v1(
qo_indptr,
kv_indptr,
kv_last_page_len,
NUM_HEADS // NUM_KV_HEADS,
NUM_KV_HEADS,
True,
work_meta_data,
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=1,
uni_seqlen_qo=1,
fast_mode=False,
max_split_per_batch=PERSISTENT_SPLITS,
intra_batch_mode=True,
dtype_q=q_dtype,
dtype_kv=kv_dtype,
)
return {
"work_meta_data": work_meta_data,
"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,
}
def _run_persistent_fp8_path(
q_view: torch.Tensor,
kv_data: dict,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
batch_size: int,
kv_seq_len: int,
sm_scale: float,
) -> torch.Tensor:
total_q = q_view.shape[0]
total_kv = batch_size * kv_seq_len
device = q_view.device
q_fp8, q_scale = _quantize_query(q_view)
kv_fp8, kv_scale = kv_data["fp8"]
kv_pages = kv_fp8.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_fp8.shape[-1])
page_ids = torch.arange(total_kv, dtype=torch.int32, device=device)
last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)
metadata = _build_persistent_metadata(
batch_size,
qo_indptr,
kv_indptr,
q_fp8.dtype,
kv_fp8.dtype,
last_page_len,
)
output = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device)
mla_decode_fwd(
q_fp8,
kv_pages,
output,
qo_indptr,
kv_indptr,
page_ids,
last_page_len,
1,
page_size=PAGE_SIZE,
nhead_kv=NUM_KV_HEADS,
sm_scale=sm_scale,
logit_cap=0.0,
num_kv_splits=PERSISTENT_SPLITS,
q_scale=q_scale,
kv_scale=kv_scale,
intra_batch_mode=True,
**metadata,
)
return output
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = int(config["batch_size"])
kv_seq_len = int(config["kv_seq_len"])
total_kv = batch_size * kv_seq_len
sm_scale = float(config.get("sm_scale", DEFAULT_SCALE))
q_view = q.view(-1, NUM_HEADS, QK_HEAD_DIM)
if _prefer_persistent_fp8_path(total_kv):
return _run_persistent_fp8_path(
q_view,
kv_data,
qo_indptr,
kv_indptr,
batch_size,
kv_seq_len,
sm_scale,
)
return _run_nonpersistent_path(
q_view,
kv_data,
qo_indptr,
kv_indptr,
batch_size,
kv_seq_len,
sm_scale,
)
scrolls · 239 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 643908.
#!POPCORN leaderboard amd-mixed-mla#!POPCORN gpu MI355X- """v13: Multi-head fused HIP MLA decode with MXFP4 KV dequant on MI355X.+ import torch- One block processes ALL 16 query heads sharing the same KV data.- KV is loaded from HBM once per block and reused for all 16 heads,- yielding ~16x bandwidth savings over the per-head v12 approach.+ from task import input_t, output_t- Two-stage flash decoding:- Stage 1: per-(batch, split) kernel. 256 threads (4 wavefronts).- All 16 Q vectors loaded to LDS (36 KB). For each KV token, dequant- MXFP4 in registers, compute dot products for all 16 heads, online- softmax update, and V accumulation -- all with KV loaded once.- Stage 2: lightweight reduce kernel merges partials from all splits- using exp-weighted combination (mathematically exact).+ 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- Bandwidth per KV token per batch element:- v12: 306 bytes x 16 heads = 4896 bytes (KV loaded 16 times)- v13: 306 bytes x 1 = 306 bytes (KV loaded once, reused)+ try:+ from aiter import per_tensor_quant_hip+ HAS_QUANT_HIP = True+ except ImportError:+ HAS_QUANT_HIP = False- No cross-call caching. Fresh output buffer each call.- """- import os-- os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"-- import torch- from torch.utils.cpp_extension import load_inline- from task import input_t, output_t-- # MLA constantsNUM_HEADS = 16NUM_KV_HEADS = 1QK_HEAD_DIM = 576V_HEAD_DIM = 512- PACKED_DIM = 288 # QK_HEAD_DIM // 2- SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)+ PAGE_SIZE = 1+ DEFAULT_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)+ FP8_DTYPE = aiter_dtypes.fp8+ LOW_PAYLOAD_TOKEN_LIMIT = 65536+ PERSISTENT_TOKEN_LIMIT = 65536+ PERSISTENT_SPLITS = 32- # ---------------------------------------------------------------------------- # HIP C++ kernel source -- multi-head fused MLA decode- # ---------------------------------------------------------------------------- HIP_SOURCE = r'''- #include <hip/hip_runtime.h>- #include <torch/extension.h>- // =========================================================================- // bf16 <-> float via bit manipulation (works on all HIP targets)- // =========================================================================- __device__ __forceinline__ float bf16_to_float(unsigned short v) {- unsigned int bits = ((unsigned int)v) << 16;- return __uint_as_float(bits);- }+ def _quantize_query(q_view: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:+ if HAS_QUANT_HIP:+ q_fp8, q_scale = per_tensor_quant_hip(q_view, quant_dtype=FP8_DTYPE)+ return q_fp8, q_scale.reshape(1)- __device__ __forceinline__ unsigned short float_to_bf16(float v) {- unsigned int bits = __float_as_uint(v);- // Round to nearest even- bits += 0x7FFF + ((bits >> 16) & 1);- return (unsigned short)(bits >> 16);- }+ finfo = torch.finfo(FP8_DTYPE)+ amax = q_view.abs().amax().clamp(min=1e-12)+ q_scale = (amax / finfo.max).to(torch.float32).reshape(1)+ q_fp8 = (q_view / q_scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)+ return q_fp8, q_scale- // =========================================================================- // FP4 E2M1 dequantization lookup table (16 entries)- // Bit layout: [sign(1) exp(2) mant(1)]- // =========================================================================- __device__ __constant__ float FP4_LUT[16] = {- 0.0f, 0.5f, 1.0f, 1.5f, 2.0f, 3.0f, 4.0f, 6.0f,- -0.0f, -0.5f, -1.0f, -1.5f, -2.0f, -3.0f, -4.0f, -6.0f- };- // =========================================================================- // E8M0 -> float: 2^(e - 127) via IEEE-754 exponent placement- // =========================================================================- __device__ __forceinline__ float e8m0_to_float(unsigned char e) {- unsigned int bits = ((unsigned int)e) << 23;- return __uint_as_float(bits);- }+ def _prefer_low_overhead_path(total_kv: int) -> bool:+ return total_kv <= LOW_PAYLOAD_TOKEN_LIMIT- // =========================================================================- // Configuration- // =========================================================================- #define QK_DIM 576- #define V_DIM 512- #define BLK_SZ 32 // MXFP4 block size (elements per E8M0 scale)- #define N_BLOCKS 18 // QK_DIM / BLK_SZ = 576/32- #define V_BLOCKS 16 // V_DIM / BLK_SZ = 512/32- #define PACK_DIM 288 // QK_DIM / 2 (bytes of packed FP4)- #define THREADS 256 // 4 wavefronts of 64- #define NHEADS 16- // =========================================================================- // Stage 1: Multi-head fused MXFP4 flash-decoding- //- // Grid: (num_splits, batch_size) <-- NOT batch_size*16!- // Block: (THREADS)- //- // Each block processes ALL 16 query heads for one batch element over one- // KV-length split. Q for all 16 heads loaded to LDS (16*576 floats = 36 KB).- // KV tokens dequantized ONCE and reused for all 16 dot products.- //- // Outputs:- // partial_out [num_splits, batch_size*16, V_DIM] float32- // partial_lse [num_splits, batch_size*16] float32- // =========================================================================+ def _prefer_persistent_fp8_path(total_kv: int) -> bool:+ return total_kv > PERSISTENT_TOKEN_LIMIT- __global__ __attribute__((amdgpu_flat_work_group_size(THREADS, THREADS)))- void mla_stage1_multihead(- const unsigned short* __restrict__ Q, // [total_q, 16, 576] bf16- const unsigned char* __restrict__ KV_packed, // [total_kv, PACK_DIM] uint8- const unsigned char* __restrict__ KV_scales, // [total_kv, N_scale_cols] uint8- float* __restrict__ partial_out, // [num_splits, total_bh, V_DIM]- float* __restrict__ partial_lse, // [num_splits, total_bh]- const int* __restrict__ kv_indptr, // [batch+1]- float sm_scale,- int total_bh,- int num_splits,- int kv_stride_scales // stride of KV_scales in dim0- ) {- int split_id = blockIdx.x;- int batch_id = blockIdx.y;- int tid = threadIdx.x;- // KV range for this batch element- int kv_start = kv_indptr[batch_id];- int kv_end = kv_indptr[batch_id + 1];- int kv_len = kv_end - kv_start;+ def _run_nonpersistent_path(+ q_view: torch.Tensor,+ kv_data: dict,+ qo_indptr: torch.Tensor,+ kv_indptr: torch.Tensor,+ batch_size: int,+ kv_seq_len: int,+ sm_scale: float,+ ) -> torch.Tensor:+ total_q = q_view.shape[0]+ total_kv = batch_size * kv_seq_len+ device = q_view.device- // Base index for this batch element's heads in the flattened bh dimension- int bh_base = batch_id * NHEADS;+ page_ids = torch.arange(total_kv, dtype=torch.int32, device=device)+ last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)+ output = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device)- if (kv_len <= 0) {- // Empty sequence: write -inf LSE and zero output for all heads- for (int h = 0; h < NHEADS; h++) {- int bh_id = bh_base + h;- if (tid == 0) {- partial_lse[split_id * total_bh + bh_id] = -1e30f;- }- int out_base = (split_id * total_bh + bh_id) * V_DIM;- for (int d = tid; d < V_DIM; d += THREADS) {- partial_out[out_base + d] = 0.0f;- }- }- return;- }+ if _prefer_low_overhead_path(total_kv):+ kv_tensor = kv_data["bf16"]+ q_input = q_view+ q_scale = None+ kv_scale = None+ else:+ kv_tensor, kv_scale = kv_data["fp8"]+ q_input, q_scale = _quantize_query(q_view)- // Split range- int split_size = (kv_len + num_splits - 1) / num_splits;- int my_start = kv_start + split_id * split_size;- int my_end = my_start + split_size;- if (my_end > kv_end) my_end = kv_end;- if (my_start >= kv_end) {- // This split has no tokens- for (int h = 0; h < NHEADS; h++) {- int bh_id = bh_base + h;- if (tid == 0) {- partial_lse[split_id * total_bh + bh_id] = -1e30f;- }- int out_base = (split_id * total_bh + bh_id) * V_DIM;- for (int d = tid; d < V_DIM; d += THREADS) {- partial_out[out_base + d] = 0.0f;- }- }- return;- }+ kv_pages = kv_tensor.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_tensor.shape[-1])- // ---- Load Q for ALL 16 heads into LDS ----- // q_lds[h][d] = Q[batch_id, h, d] converted to float- // 16 * 576 = 9216 floats = 36,864 bytes (fits in LDS)- __shared__ float q_lds[NHEADS][QK_DIM];-- int q_token = batch_id; // In decode, total_q == batch_size- for (int h = 0; h < NHEADS; h++) {- long long q_base = (long long)q_token * NHEADS * QK_DIM + (long long)h * QK_DIM;- for (int d = tid; d < QK_DIM; d += THREADS) {- q_lds[h][d] = bf16_to_float(Q[q_base + d]);- }- }- __syncthreads();-- // ---- Per-head online softmax state (in registers) ----- float my_max[NHEADS];- float my_sum_exp[NHEADS];- float my_v0[NHEADS]; // V accumulator for dim = tid- float my_v1[NHEADS]; // V accumulator for dim = tid + 256-- for (int h = 0; h < NHEADS; h++) {- my_max[h] = -1e30f;- my_sum_exp[h] = 0.0f;- my_v0[h] = 0.0f;- my_v1[h] = 0.0f;- }-- // ---- Shared memory for cross-warp reduction ----- // warp_dots[warp_id][head] for dot product reduction- // We process heads in groups to limit shared memory- __shared__ float warp_dots[4][NHEADS];- // For broadcasting final scores to all threads- __shared__ float score_broadcast[NHEADS];-- int warp_id = tid / 64;- int lane_id = tid % 64;-- // ---- Iterate over KV tokens in this split ----- for (int kv_pos = my_start; kv_pos < my_end; kv_pos++) {-- long long packed_base = (long long)kv_pos * PACK_DIM;- long long scale_base = (long long)kv_pos * kv_stride_scales;-- // -- Step 1: Dequantize KV values for this thread's dimensions --- // Each thread handles dims: tid, tid+256, and tid+512 (if tid < 64)- // We dequant once and reuse for all 16 heads-- // Dequant dim = tid (always valid, tid < 256 < 576)- int d0 = tid;- int byte_idx0 = d0 / 2;- unsigned char pk0 = KV_packed[packed_base + byte_idx0];- float rv0;- if (d0 & 1) rv0 = FP4_LUT[(pk0 >> 4) & 0x0F];- else rv0 = FP4_LUT[pk0 & 0x0F];- float sc0 = e8m0_to_float(KV_scales[scale_base + d0 / BLK_SZ]);- float kv_d0 = rv0 * sc0;-- // Dequant dim = tid + 256 (always valid, tid+256 < 512 < 576)- int d1 = tid + THREADS;- int byte_idx1 = d1 / 2;- unsigned char pk1 = KV_packed[packed_base + byte_idx1];- float rv1;- if (d1 & 1) rv1 = FP4_LUT[(pk1 >> 4) & 0x0F];- else rv1 = FP4_LUT[pk1 & 0x0F];- float sc1 = e8m0_to_float(KV_scales[scale_base + d1 / BLK_SZ]);- float kv_d1 = rv1 * sc1;-- // Dequant dim = tid + 512 (only valid for tid < 64, since 576-512=64)- float kv_d2 = 0.0f;- if (tid < 64) {- int d2 = tid + 2 * THREADS;- int byte_idx2 = d2 / 2;- unsigned char pk2 = KV_packed[packed_base + byte_idx2];- float rv2;- if (d2 & 1) rv2 = FP4_LUT[(pk2 >> 4) & 0x0F];- else rv2 = FP4_LUT[pk2 & 0x0F];- float sc2 = e8m0_to_float(KV_scales[scale_base + d2 / BLK_SZ]);- kv_d2 = rv2 * sc2;- }-- // -- Step 2: Compute dot products for ALL 16 heads --- // Each thread computes partial dot for each head using its KV dims- float pdot[NHEADS];- for (int h = 0; h < NHEADS; h++) {- pdot[h] = q_lds[h][d0] * kv_d0- + q_lds[h][d1] * kv_d1;- if (tid < 64) {- pdot[h] += q_lds[h][tid + 2 * THREADS] * kv_d2;- }- }-- // -- Step 3: Warp-level reduction for each head --- for (int h = 0; h < NHEADS; h++) {- float val = pdot[h];- for (int offset = 32; offset >= 1; offset >>= 1) {- val += __shfl_xor(val, offset);- }- // Lane 0 of each warp has the warp sum- if (lane_id == 0) {- warp_dots[warp_id][h] = val;- }- }- __syncthreads();-- // -- Step 4: Thread 0 does final cross-warp reduction for all heads --- if (tid == 0) {- for (int h = 0; h < NHEADS; h++) {- float s = warp_dots[0][h] + warp_dots[1][h]- + warp_dots[2][h] + warp_dots[3][h];- score_broadcast[h] = s * sm_scale;- }- }- __syncthreads();-- // -- Step 5: Online softmax update and V accumulation for all heads --- // V = first 512 dims of KV. Each thread owns V[tid] and V[tid+256].- // kv_d0 is the dequanted value at dim=tid (always < 512, valid V dim)- // kv_d1 is the dequanted value at dim=tid+256 (always < 512, valid V dim)- float v_val0 = kv_d0; // V dim = tid- float v_val1 = kv_d1; // V dim = tid + 256-- for (int h = 0; h < NHEADS; h++) {- float score = score_broadcast[h];-- float old_max = my_max[h];- my_max[h] = fmaxf(my_max[h], score);- float corr = expf(old_max - my_max[h]);- float w = expf(score - my_max[h]);- my_sum_exp[h] = my_sum_exp[h] * corr + w;-- my_v0[h] = my_v0[h] * corr + w * v_val0;- my_v1[h] = my_v1[h] * corr + w * v_val1;- }-- __syncthreads(); // needed before next iteration reuses shared mem- }-- // ---- Write partial results for all 16 heads ----- for (int h = 0; h < NHEADS; h++) {- int bh_id = bh_base + h;- int out_base = (split_id * total_bh + bh_id) * V_DIM;-- float inv_se = (my_sum_exp[h] > 0.0f) ? (1.0f / my_sum_exp[h]) : 0.0f;- partial_out[out_base + tid] = my_v0[h] * inv_se;- partial_out[out_base + tid + THREADS] = my_v1[h] * inv_se;-- if (tid == 0) {- float lse_val = (my_sum_exp[h] > 0.0f)- ? (my_max[h] + logf(my_sum_exp[h]))- : -1e30f;- partial_lse[split_id * total_bh + bh_id] = lse_val;- }- }- }--- // =========================================================================- // Stage 2: Reduce partials across splits- //- // Grid: (total_bh)- // Block: (THREADS)- //- // For each (batch, head), merges all split partials into final bf16 output.- // (Same as v12 -- this is already per-head and efficient.)- // =========================================================================-- __global__ __attribute__((amdgpu_flat_work_group_size(THREADS, THREADS)))- void mla_reduce(- const float* __restrict__ partial_out, // [num_splits, total_bh, V_DIM]- const float* __restrict__ partial_lse, // [num_splits, total_bh]- unsigned short* __restrict__ final_out, // [total_bh, V_DIM] bf16- int total_bh,- int num_splits- ) {- int bh_id = blockIdx.x;- int tid = threadIdx.x;-- // Pass 1: find global max LSE- __shared__ float gmax_shared;- if (tid == 0) {- float gmax = -1e30f;- for (int s = 0; s < num_splits; s++) {- float lse_s = partial_lse[s * total_bh + bh_id];- if (lse_s > gmax) gmax = lse_s;- }- gmax_shared = gmax;- }- __syncthreads();- float gmax = gmax_shared;-- // Pass 2: compute total weight- __shared__ float total_w_shared;- if (tid == 0) {- float total_w = 0.0f;- for (int s = 0; s < num_splits; s++) {- float lse_s = partial_lse[s * total_bh + bh_id];- total_w += expf(lse_s - gmax);- }- total_w_shared = total_w;- }- __syncthreads();- float total_w = total_w_shared;- float inv_tw = (total_w > 0.0f) ? (1.0f / total_w) : 0.0f;-- // Pass 3: weighted merge of V values- float acc0 = 0.0f;- float acc1 = 0.0f;-- for (int s = 0; s < num_splits; s++) {- float lse_s = partial_lse[s * total_bh + bh_id];- float w_s = expf(lse_s - gmax);- int base = (s * total_bh + bh_id) * V_DIM;- acc0 += partial_out[base + tid] * w_s;- acc1 += partial_out[base + tid + THREADS] * w_s;- }-- // Write bf16 output- long long out_base = (long long)bh_id * V_DIM;- final_out[out_base + tid] = float_to_bf16(acc0 * inv_tw);- final_out[out_base + tid + THREADS] = float_to_bf16(acc1 * inv_tw);- }--- // =========================================================================- // C++ / pybind11 wrapper- // =========================================================================-- torch::Tensor mla_decode_hip(- torch::Tensor Q, // [total_q, 16, 576] bf16- torch::Tensor KV_packed, // [total_kv, PACK_DIM] uint8- torch::Tensor KV_scales, // [total_kv, N_scale_cols] uint8- torch::Tensor kv_indptr, // [batch+1] int32- int64_t batch_size,- int64_t num_splits,- double sm_scale_d- ) {- float sm_scale = (float)sm_scale_d;- int total_q = Q.size(0);- int total_bh = total_q * NHEADS;-- // KV scale stride (column count, may be padded)- int kv_stride_scales = KV_scales.size(1);-- // Allocate partial buffers- auto opts_f32 = torch::TensorOptions().dtype(torch::kFloat32).device(Q.device());- auto partial_out = torch::zeros({(int64_t)num_splits, (int64_t)total_bh, (int64_t)V_DIM}, opts_f32);- auto partial_lse = torch::full({(int64_t)num_splits, (int64_t)total_bh}, -1e30f, opts_f32);-- // Stage 1: multi-head fused kernel -- grid is (splits, batch) not (splits, batch*heads)- dim3 grid_s1((int)num_splits, (int)batch_size);- dim3 block_s1(THREADS);- mla_stage1_multihead<<<grid_s1, block_s1, 0, 0>>>(- (const unsigned short*)Q.data_ptr(),- (const unsigned char*)KV_packed.data_ptr(),- (const unsigned char*)KV_scales.data_ptr(),- partial_out.data_ptr<float>(),- partial_lse.data_ptr<float>(),- kv_indptr.data_ptr<int>(),- sm_scale,- total_bh,- (int)num_splits,- kv_stride_scales- );-- // Stage 2: reduce across splits -> bf16 output- auto opts_bf16 = torch::TensorOptions().dtype(torch::kBFloat16).device(Q.device());- auto output = torch::empty({(int64_t)total_q, NHEADS, (int64_t)V_DIM}, opts_bf16);-- dim3 grid_s2(total_bh);- dim3 block_s2(THREADS);- mla_reduce<<<grid_s2, block_s2, 0, 0>>>(- partial_out.data_ptr<float>(),- partial_lse.data_ptr<float>(),- (unsigned short*)output.data_ptr(),- total_bh,- (int)num_splits- );-- return output;- }-- PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {- m.def("mla_decode_hip", &mla_decode_hip,- "MLA decode with multi-head fused MXFP4 dequant (HIP kernel)");- }- '''-- # ---------------------------------------------------------------------------- # Compile HIP kernel- # ----------------------------------------------------------------------------- import sys-- print("[mla_v13] Compiling multi-head fused HIP MXFP4 MLA decode kernel...", file=sys.stderr, flush=True)- try:- _mod = load_inline(- name="mla_v13_hip",- cpp_sources=[],- cuda_sources=[HIP_SOURCE],- extra_cuda_cflags=["-O3", "--offload-arch=gfx950"],- verbose=False,+ mla_decode_fwd(+ q_input,+ kv_pages,+ output,+ qo_indptr,+ kv_indptr,+ page_ids,+ last_page_len,+ 1,+ page_size=PAGE_SIZE,+ nhead_kv=NUM_KV_HEADS,+ sm_scale=sm_scale,+ q_scale=q_scale,+ kv_scale=kv_scale,+ intra_batch_mode=False,)- _HIP_AVAILABLE = True- print("[mla_v13] HIP compilation successful!", file=sys.stderr, flush=True)- except Exception as e:- _HIP_AVAILABLE = False- print(f"[mla_v13] HIP compilation failed: {e}", file=sys.stderr, flush=True)- # ---------------------------------------------------------------------------- # AITER fallback imports (fp8 path)- # ---------------------------------------------------------------------------+ return output- from aiter.mla import mla_decode_fwd- from aiter import dtypes as aiter_dtypes, get_mla_metadata_info_v1, get_mla_metadata_v1- FP8 = aiter_dtypes.fp8- PAGE_SIZE = 1--- def _quantize_fp8(t):- finfo = torch.finfo(FP8)- amax = t.abs().amax().clamp(min=1e-12)- sc = amax / finfo.max- return (t / sc).clamp(finfo.min, finfo.max).to(FP8), sc.float().reshape(1)--- def _choose_splits(total_kv, kv_seq_len):- """Select number of KV splits for the HIP kernel."""- if kv_seq_len <= 512:- return 2- elif kv_seq_len <= 2048:- return 4- elif kv_seq_len <= 8192:- return 8- elif kv_seq_len <= 32768:- return 16- else:- return 32--- def _choose_aiter_params(qsl, total_kv, bs, nq):- kv_per_req = total_kv // max(bs, 1)- if qsl == 1:- if kv_per_req >= 4096:- return 32, False- elif kv_per_req >= 1024:- return 16, False- else:- return 8, False- else:- if kv_per_req >= 4096:- return 16, False- elif kv_per_req >= 1024:- return 8, False- else:- return 4, False--- def _build_meta(bs, qsl, nq, nkv, total_kv, q_dtype, kv_dtype,- qo_indptr, kv_indptr, num_kv_splits, fast_mode):- kv_lpl = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)+ def _build_persistent_metadata(+ batch_size: int,+ qo_indptr: torch.Tensor,+ kv_indptr: torch.Tensor,+ q_dtype: torch.dtype,+ kv_dtype: torch.dtype,+ kv_last_page_len: torch.Tensor,+ ):info = get_mla_metadata_info_v1(- bs, qsl, nq, q_dtype, kv_dtype,- is_sparse=False, fast_mode=fast_mode,- num_kv_splits=num_kv_splits, intra_batch_mode=True,+ batch_size,+ 1,+ NUM_HEADS,+ q_dtype,+ kv_dtype,+ is_sparse=False,+ fast_mode=False,+ num_kv_splits=PERSISTENT_SPLITS,+ intra_batch_mode=True,)- work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]- wm, wi, wis, ri, rfm, rpm = work++ work_meta_data, work_indptr, work_info_set, reduce_indptr, reduce_final_map, reduce_partial_map = [+ torch.empty(shape, dtype=dtype, device=qo_indptr.device) for shape, dtype in info+ ]+get_mla_metadata_v1(- qo_indptr, kv_indptr, kv_lpl,- nq // nkv, nkv, True,- wm, wis, wi, ri, rfm, rpm,- page_size=PAGE_SIZE, kv_granularity=max(PAGE_SIZE, 16),- max_seqlen_qo=qsl, uni_seqlen_qo=qsl,- fast_mode=fast_mode, max_split_per_batch=num_kv_splits,- intra_batch_mode=True, dtype_q=q_dtype, dtype_kv=kv_dtype,+ qo_indptr,+ kv_indptr,+ kv_last_page_len,+ NUM_HEADS // NUM_KV_HEADS,+ NUM_KV_HEADS,+ True,+ work_meta_data,+ 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=1,+ uni_seqlen_qo=1,+ fast_mode=False,+ max_split_per_batch=PERSISTENT_SPLITS,+ intra_batch_mode=True,+ dtype_q=q_dtype,+ dtype_kv=kv_dtype,)+return {- "meta": dict(work_meta_data=wm, work_indptr=wi, work_info_set=wis,- reduce_indptr=ri, reduce_final_map=rfm, reduce_partial_map=rpm),- "kv_lpl": kv_lpl,- "kv_indices": torch.arange(total_kv, dtype=torch.int32, device="cuda"),- "num_kv_splits": num_kv_splits,+ "work_meta_data": work_meta_data,+ "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,}- def _run_mla_aiter(q, kv_data, qo_indptr, kv_indptr, config):- """AITER fp8 fallback path."""- nq = config["num_heads"]- nkv = config["num_kv_heads"]- dq = config["qk_head_dim"]- dv = config["v_head_dim"]- qsl = config["q_seq_len"]- bs = config["batch_size"]+ def _run_persistent_fp8_path(+ q_view: torch.Tensor,+ kv_data: dict,+ qo_indptr: torch.Tensor,+ kv_indptr: torch.Tensor,+ batch_size: int,+ kv_seq_len: int,+ sm_scale: float,+ ) -> torch.Tensor:+ total_q = q_view.shape[0]+ total_kv = batch_size * kv_seq_len+ device = q_view.device+ q_fp8, q_scale = _quantize_query(q_view)kv_fp8, kv_scale = kv_data["fp8"]- total_kv = int(kv_indptr[-1].item())+ kv_pages = kv_fp8.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_fp8.shape[-1])- q_fp8, q_scale = _quantize_fp8(q)+ page_ids = torch.arange(total_kv, dtype=torch.int32, device=device)+ last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device=device)+ metadata = _build_persistent_metadata(+ batch_size,+ qo_indptr,+ kv_indptr,+ q_fp8.dtype,+ kv_fp8.dtype,+ last_page_len,+ )- num_kv_splits, fast_mode = _choose_aiter_params(qsl, total_kv, bs, nq)- cached = _build_meta(bs, qsl, nq, nkv, total_kv,- q_fp8.dtype, kv_fp8.dtype,- qo_indptr, kv_indptr, num_kv_splits, fast_mode)+ output = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=device)- kv_4d = kv_fp8.view(kv_fp8.shape[0], PAGE_SIZE, nkv, kv_fp8.shape[-1])- o = torch.empty((q_fp8.shape[0], nq, dv), dtype=torch.bfloat16, device="cuda")-mla_decode_fwd(- q_fp8.view(-1, nq, dq), kv_4d, o,- qo_indptr, kv_indptr,- cached["kv_indices"], cached["kv_lpl"], qsl,- page_size=PAGE_SIZE, nhead_kv=nkv,- sm_scale=SM_SCALE, logit_cap=0.0,- num_kv_splits=cached["num_kv_splits"],- q_scale=q_scale, kv_scale=kv_scale,- intra_batch_mode=True, **cached["meta"],+ q_fp8,+ kv_pages,+ output,+ qo_indptr,+ kv_indptr,+ page_ids,+ last_page_len,+ 1,+ page_size=PAGE_SIZE,+ nhead_kv=NUM_KV_HEADS,+ sm_scale=sm_scale,+ logit_cap=0.0,+ num_kv_splits=PERSISTENT_SPLITS,+ q_scale=q_scale,+ kv_scale=kv_scale,+ intra_batch_mode=True,+ **metadata,)- return o+ return output- # ---------------------------------------------------------------------------- # Entry point- # ---------------------------------------------------------------------------def custom_kernel(data: input_t) -> output_t:- """MLA decode with multi-head fused MXFP4 HIP kernel (fallback: AITER fp8)."""q, kv_data, qo_indptr, kv_indptr, config = data- batch_size = config["batch_size"]- kv_seq_len = config["kv_seq_len"]- qsl = config["q_seq_len"]+ batch_size = int(config["batch_size"])+ kv_seq_len = int(config["kv_seq_len"])+ total_kv = batch_size * kv_seq_len+ sm_scale = float(config.get("sm_scale", DEFAULT_SCALE))+ q_view = q.view(-1, NUM_HEADS, QK_HEAD_DIM)- # ---- Try HIP kernel path (MXFP4 dequant, multi-head fusion) ----- if _HIP_AVAILABLE and qsl == 1:- kv_fp4x2, kv_scales_e8m0 = kv_data["mxfp4"]+ if _prefer_persistent_fp8_path(total_kv):+ return _run_persistent_fp8_path(+ q_view,+ kv_data,+ qo_indptr,+ kv_indptr,+ batch_size,+ kv_seq_len,+ sm_scale,+ )- # Reshape to 2D: [total_kv, PACK_DIM] uint8- kv_packed_u8 = kv_fp4x2.reshape(-1, PACKED_DIM).contiguous().view(torch.uint8)-- # Scales: [total_kv_padded, N_cols] uint8 (N_cols >= 18, may be padded)- kv_scales_u8 = kv_scales_e8m0.contiguous().view(torch.uint8)-- # Ensure scales are 2D [total_kv, N_cols]- total_kv = int(kv_indptr[-1].item())- if kv_scales_u8.dim() == 1:- n_scale_cols = kv_scales_u8.numel() // total_kv- kv_scales_u8 = kv_scales_u8.view(total_kv, n_scale_cols)- elif kv_scales_u8.dim() != 2:- kv_scales_u8 = kv_scales_u8.view(total_kv, -1)-- # If padded (more rows than total_kv), just use total_kv rows- if kv_scales_u8.size(0) > total_kv:- kv_scales_u8 = kv_scales_u8[:total_kv]-- num_splits = _choose_splits(total_kv, kv_seq_len)-- try:- output = _mod.mla_decode_hip(- q, kv_packed_u8, kv_scales_u8,- kv_indptr, batch_size, num_splits,- float(SM_SCALE),- )- return output- except Exception:- # Fall through to AITER- pass-- # ---- AITER fp8 fallback ----- return _run_mla_aiter(q, kv_data, qo_indptr, kv_indptr, config)+ return _run_nonpersistent_path(+ q_view,+ kv_data,+ qo_indptr,+ kv_indptr,+ batch_size,+ kv_seq_len,+ sm_scale,+ )
scrolls · 822 diff lines total
Best evidence level for this revision: reported
JSON