submission 643908
Maxwell Cipher · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 637 lines, June 9 Researcher Reciprocity License v1.0.
mla_v13.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-643908?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:87aff7105ee8fbcc1b844f5de8032e1191945924391dee27a75bb87a7322e407
license declaredunknown
license concludedunknown
authorsMaxwell Cipher
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
"""v13: Multi-head fused HIP MLA decode with MXFP4 KV dequant on MI355X.shared-memory
__shared__ float q_lds[NHEADS][QK_DIM];Kernel source
mla_v13.py637 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""v13: Multi-head fused HIP MLA decode with MXFP4 KV dequant on MI355X.
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.
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).
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)
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 constants
NUM_HEADS = 16
NUM_KV_HEADS = 1
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
PACKED_DIM = 288 # QK_HEAD_DIM // 2
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
# ---------------------------------------------------------------------------
# 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);
}
__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);
}
// =========================================================================
// 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);
}
// =========================================================================
// 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
// =========================================================================
__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;
// Base index for this batch element's heads in the flattened bh dimension
int bh_base = batch_id * NHEADS;
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;
}
// 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;
}
// ---- 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,
)
_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)
# ---------------------------------------------------------------------------
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)
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,
)
work = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
wm, wi, wis, ri, rfm, rpm = work
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,
)
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,
}
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"]
kv_fp8, kv_scale = kv_data["fp8"]
total_kv = int(kv_indptr[-1].item())
q_fp8, q_scale = _quantize_fp8(q)
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)
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"],
)
return o
# ---------------------------------------------------------------------------
# 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"]
# ---- Try HIP kernel path (MXFP4 dequant, multi-head fusion) ----
if _HIP_AVAILABLE and qsl == 1:
kv_fp4x2, kv_scales_e8m0 = kv_data["mxfp4"]
# 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)
scrolls · 637 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