submission 590963
gwokhou · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1119 lines, June 9 Researcher Reciprocity License v1.0.
submission_v4.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-590963?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:a61fb8eccef7ec61087a9e95187900937f869f1e02e49cb8ec56445dbb659129
license declaredunknown
license concludedunknown
authorsgwokhou
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Submission v4: inlines mla_decode_fwd and MXFP4 stage1 from aiter (no import of mla_decode_fwd).num-warps = 4
num_warps=4,persistent-kernel
Decode only — persistent mode with get_mla_metadata_v1.split-k
for split_kv_id in range(0, num_valid_kv_splits):stages = 2
num_stages=2,tile-k = 32
BLOCK_K = 32tile-n = 64
BLOCK_N = 64Kernel source
submission_v4.py1119 lines
"""
Reference implementation for MLA (Multi-head Latent Attention) decode kernel.
Submission v4: inlines mla_decode_fwd and MXFP4 stage1 from aiter (no import of mla_decode_fwd).
Uses the same aiter MLA API; mla_decode_fwd and its children call chain
(get_meta_param, _fwd_kernel_stage2_asm, mla_decode_stage1_mxfp4) are copied here.
DeepSeek R1 forward_absorb MLA: absorbed q (576), compressed kv_buffer (576),
output v_head_dim = kv_lora_rank = 512.
Decode only — persistent mode with get_mla_metadata_v1.
"""
from __future__ import annotations
import functools
import aiter
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.jit.utils.chip_info import get_cu_num, get_gfx
from aiter.ops.triton.utils._triton import arch_info
from aiter.utility.fp4_utils import dynamic_mxfp4_quant, e8m0_to_f32, mxfp4_to_f32
from task import input_t, output_t
import torch
import triton
import triton.language as tl
# ---------------------------------------------------------------------------
# DeepSeek R1 latent MQA constants (forward_absorb path)
# ---------------------------------------------------------------------------
TOTAL_NUM_HEADS = 128
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM # 576
V_HEAD_DIM = KV_LORA_RANK # 512
SM_SCALE = 1.0 / (QK_HEAD_DIM**0.5)
PAGE_SIZE = 1
NUM_KV_SPLITS = 32
FP8_DTYPE = aiter_dtypes.fp8
MXFP4_DTYPE = aiter_dtypes.fp4x2
QKV_DTYPE = "mxfp4"
# MXFP4 block size (must match fp4_utils.dynamic_mxfp4_quant)
MXFP4_BLOCK_SIZE = 32
# ---------------------------------------------------------------------------
# FP8 quantization
# ---------------------------------------------------------------------------
def quantize_fp8(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""
Dynamic per-tensor FP8 quantization (following sglang scaled_fp8_quant).
Args:
tensor: bf16 tensor to quantize
Returns:
(fp8_tensor, scale) where scale is a scalar float32 tensor.
Dequantize: fp8_tensor.to(bf16) * scale
"""
finfo = torch.finfo(FP8_DTYPE)
amax = tensor.abs().amax().clamp(min=1e-12)
scale = amax / finfo.max
fp8_tensor = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
return fp8_tensor, scale.to(torch.float32).reshape(1)
# ---------------------------------------------------------------------------
# MXFP4 quantization (aiter native: block-32)
# ---------------------------------------------------------------------------
def quantize_mxfp4(tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
orig_shape = tensor.shape
B, M, N = orig_shape
tensor_2d = tensor.reshape(B * M, N)
fp4_data_2d, scale_e8m0 = dynamic_mxfp4_quant(tensor_2d)
fp4_data = fp4_data_2d.view(B, M, N // 2)
return fp4_data, scale_e8m0
def dequantize_mxfp4(
fp4_data: torch.Tensor,
scale_e8m0: torch.Tensor,
orig_shape: tuple,
dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor:
"""
Dequantize MXFP4 tensor using aiter utilities.
Note: dynamic_mxfp4_quant may pad both row and block dimensions in scale_e8m0.
We trim scales to match the actual data dimensions.
Args:
fp4_data: packed FP4 data, shape [B, M, N//2] in fp4x2 or uint8
scale_e8m0: E8M0 block scale factors (possibly padded) in fp8_e8m0
orig_shape: original (B, M, N) for reshaping
dtype: output dtype
Returns:
Dequantized tensor of shape orig_shape.
"""
B, M, N = orig_shape
num_rows = B * M
block_size = 32
num_blocks = N // block_size
fp4_data_2d = fp4_data.reshape(num_rows, N // 2)
float_vals = mxfp4_to_f32(fp4_data_2d)
scale_f32 = e8m0_to_f32(scale_e8m0)
scale_f32 = scale_f32[:num_rows, :num_blocks]
float_vals_blocked = float_vals.view(num_rows, num_blocks, block_size)
scaled = float_vals_blocked * scale_f32.unsqueeze(-1)
return scaled.view(B, M, N).to(dtype)
# ---------------------------------------------------------------------------
# Inlined from aiter.mla: stage2 kernel + get_meta_param + mla_decode_fwd
# ---------------------------------------------------------------------------
@triton.jit
def _fwd_kernel_stage2_asm(
Mid_O,
Mid_lse,
O,
qo_indptr,
kv_indptr,
num_kv_splits_indptr,
stride_mid_ob: tl.int64,
stride_mid_oh: tl.int64,
stride_mid_os: tl.int64,
stride_obs: tl.int64,
stride_oh: tl.int64,
MAYBE_FINAL_OUT: tl.constexpr,
BATCH_NUM: tl.constexpr,
BLOCK_DV: tl.constexpr,
Lv: tl.constexpr,
mgc: tl.constexpr,
):
cur_batch = tl.program_id(0)
cur_head = tl.program_id(1)
cur_qo_start = tl.load(qo_indptr + cur_batch)
cur_qo_end = tl.load(qo_indptr + cur_batch + 1)
cur_split_start = tl.load(num_kv_splits_indptr + cur_batch)
cur_split_end = tl.load(num_kv_splits_indptr + cur_batch + 1)
num_max_kv_splits = tl.load(num_kv_splits_indptr + BATCH_NUM)
cur_kv_seq_len = tl.load(kv_indptr + cur_batch + 1) - tl.load(kv_indptr + cur_batch)
offs_d = tl.arange(0, BLOCK_DV)
mask_d = offs_d < Lv
offs_logic = cur_qo_start * stride_mid_ob + cur_head * stride_mid_oh
offs_v = offs_logic * Lv + offs_d
num_valid_kv_splits = tl.minimum(
cur_split_end - cur_split_start, tl.cdiv(cur_kv_seq_len, mgc)
)
FINAL_OUT = MAYBE_FINAL_OUT and num_max_kv_splits == BATCH_NUM
for cur_qo in range(cur_qo_start, cur_qo_end):
if FINAL_OUT:
input_ptr = Mid_O.to(tl.pointer_type(O.type.element_ty))
out = tl.load(
input_ptr
+ Lv * (cur_qo * stride_mid_os + cur_head * stride_mid_oh)
+ offs_d,
mask=mask_d,
other=0.0,
)
tl.store(
O + cur_qo * stride_obs + cur_head * stride_oh + offs_d,
out,
mask=mask_d,
)
else:
e_sum = 0.0
e_max = -float("inf")
acc = tl.zeros((BLOCK_DV,), dtype=tl.float32)
for split_kv_id in range(0, num_valid_kv_splits):
tv = tl.load(
Mid_O + offs_v + split_kv_id * stride_mid_os * Lv,
mask=mask_d,
other=0.0,
)
tlogic = tl.load(Mid_lse + offs_logic + split_kv_id * stride_mid_os)
n_e_max = tl.maximum(tlogic, e_max)
old_scale = tl.exp(e_max - n_e_max)
acc *= old_scale
exp_logic = tl.exp(tlogic - n_e_max)
acc += exp_logic * tv
e_sum = e_sum * old_scale + exp_logic
e_max = n_e_max
offs_logic += stride_mid_ob
offs_v += stride_mid_ob * Lv
tl.store(
O + cur_qo * stride_obs + cur_head * stride_oh + offs_d,
acc / e_sum,
mask=mask_d,
)
@functools.lru_cache()
def get_meta_param(num_kv_splits, bs, total_kv, nhead, max_seqlen_q, dtype):
if num_kv_splits is None:
cu_num = get_cu_num()
avg_kv = total_kv / bs
overhead = 84.1
tmp = [
(
bs
* i
/ ((bs * i + cu_num - 1) // cu_num * cu_num)
* avg_kv
/ (avg_kv + overhead * i),
i,
)
for i in range(1, 17)
]
num_kv_splits = sorted(tmp, key=lambda x: x[0], reverse=True)[0][1]
get_block_n_fp8 = {
16: 128,
32: 128,
48: 64,
64: 64,
128: 32,
256: 32,
384: 32,
512: 32,
}
if dtype == aiter_dtypes.fp8:
min_block_n = get_block_n_fp8[int(nhead * max_seqlen_q)]
num_kv_splits = min(
num_kv_splits, int(total_kv / bs + min_block_n - 1) // min_block_n
)
num_kv_splits_indptr = torch.arange(
0, (bs + 1) * num_kv_splits, num_kv_splits, dtype=torch.int, device="cuda"
)
return num_kv_splits, num_kv_splits_indptr
def mla_decode_fwd(
q,
kv_buffer,
o,
qo_indptr,
kv_indptr,
kv_indices,
kv_last_page_lens,
max_seqlen_q,
page_size=1,
nhead_kv=1,
sm_scale=None,
logit_cap=0.0,
num_kv_splits=None,
num_kv_splits_indptr=None,
work_meta_data=None,
work_indptr=None,
work_info_set=None,
reduce_indptr=None,
reduce_final_map=None,
reduce_partial_map=None,
q_scale=None,
kv_scale=None,
intra_batch_mode=False,
return_logits=False,
return_lse=False,
use_mxfp4=False,
):
device = q.device
assert logit_cap <= 0, f"{logit_cap=} is not support yet"
if kv_buffer.dtype != torch.uint8:
_, _, _, qk_head_dim = kv_buffer.shape
else:
_, _, qk_head_dim = q.shape
if sm_scale is None:
sm_scale = 1.0 / (qk_head_dim**0.5)
ori_total_s, ori_nhead, ori_v_head_dim = o.shape
total_s, nhead, v_head_dim = o.shape
bs = qo_indptr.shape[0] - 1
total_kv = kv_indices.shape[0]
persistent_mode = work_meta_data is not None
io_transformed = False
if not persistent_mode:
if num_kv_splits is None or num_kv_splits_indptr is None:
num_kv_splits, num_kv_splits_indptr = get_meta_param(
num_kv_splits, bs, total_kv, nhead, max_seqlen_q, q.dtype
)
mgc = 64 if max_seqlen_q == 1 and nhead == 16 else 16
MAYBE_FINAL_OUT = True
if nhead == 16 and max_seqlen_q == 1:
MAYBE_FINAL_OUT = False
logits = (
o.view((total_s, num_kv_splits, nhead, v_head_dim))
if (
num_kv_splits == 1
and (
q.dtype == aiter_dtypes.fp8
or (q.dtype == aiter_dtypes.bf16 and max_seqlen_q == 4)
)
)
else torch.empty(
(total_s, num_kv_splits, nhead, v_head_dim),
dtype=torch.float32,
device=device,
)
)
attn_lse = torch.empty(
(total_s, num_kv_splits, nhead, 1), dtype=torch.float32, device=device
)
final_lse = torch.empty((total_s, nhead), dtype=torch.float32, device=device)
use_mxfp4_path = (
use_mxfp4
and get_gfx() == "gfx950"
and arch_info.is_fp4_avail()
and q.dtype == aiter_dtypes.bf16
and kv_buffer.dtype == aiter_dtypes.bf16
and qk_head_dim % 32 == 0
and v_head_dim % 32 == 0
)
if use_mxfp4_path:
mla_decode_stage1_mxfp4(
q,
kv_buffer,
qo_indptr,
kv_indptr,
kv_indices,
kv_last_page_lens,
num_kv_splits_indptr,
max_seqlen_q,
page_size,
nhead_kv,
sm_scale,
logits,
attn_lse,
o,
)
else:
aiter.mla_decode_stage1_asm_fwd(
q,
kv_buffer,
qo_indptr,
kv_indptr,
kv_indices,
kv_last_page_lens,
num_kv_splits_indptr,
None,
None,
None,
max_seqlen_q,
page_size,
nhead_kv,
sm_scale,
logits,
attn_lse,
o,
q_scale,
kv_scale,
)
if num_kv_splits == 1 and (
use_mxfp4_path
or q.dtype == aiter_dtypes.fp8
or (q.dtype == aiter_dtypes.bf16 and max_seqlen_q == 4)
or (
q.dtype == aiter_dtypes.bf16
and kv_buffer.dtype == aiter_dtypes.bf16
and nhead in [32, 64]
)
):
return logits.view(total_s, nhead, v_head_dim), attn_lse
Lv = v_head_dim
BLOCK_DV = triton.next_power_of_2(Lv)
grid = (bs, nhead)
extra_kargs = {"waves_per_eu": 4}
_fwd_kernel_stage2_asm[grid](
logits,
attn_lse,
o,
qo_indptr,
kv_indptr,
num_kv_splits_indptr,
attn_lse.stride(0),
attn_lse.stride(2),
attn_lse.stride(1),
o.stride(0),
o.stride(1),
MAYBE_FINAL_OUT=MAYBE_FINAL_OUT,
BATCH_NUM=bs,
BLOCK_DV=BLOCK_DV,
Lv=Lv,
mgc=mgc,
num_warps=4,
num_stages=2,
**extra_kargs,
)
else:
if num_kv_splits is None:
num_kv_splits = get_cu_num()
if (
nhead == 16
or (
nhead == 128
and q.dtype == aiter_dtypes.fp8
and kv_buffer.dtype == aiter_dtypes.fp8
)
or (
get_gfx() == "gfx950"
and nhead == 32
and q.dtype == aiter_dtypes.fp8
and kv_buffer.dtype == aiter_dtypes.fp8
and max_seqlen_q == 4
)
):
pass
elif nhead in range(32, 128 + 1, 16) and persistent_mode:
total_s = ori_total_s * (ori_nhead // 16)
nhead = 16
q = q.view(total_s, nhead, -1)
o = o.view(total_s, nhead, -1)
io_transformed = True
else:
assert False, f"{nhead=} and {max_seqlen_q=} not supported"
logits = torch.empty(
(reduce_partial_map.size(0) * max_seqlen_q, 1, nhead, v_head_dim),
dtype=torch.float32,
device=device,
)
attn_lse = torch.empty(
(reduce_partial_map.size(0) * max_seqlen_q, 1, nhead, 1),
dtype=torch.float32,
device=device,
)
final_lse = (
torch.empty((total_s, nhead), dtype=torch.float32, device=device)
if return_lse
else None
)
aiter.mla_decode_stage1_asm_fwd(
q,
kv_buffer,
qo_indptr,
kv_indptr,
kv_indices,
kv_last_page_lens,
num_kv_splits_indptr,
work_meta_data,
work_indptr,
work_info_set,
max_seqlen_q,
page_size,
nhead_kv,
sm_scale,
logits,
attn_lse,
o,
q_scale,
kv_scale,
)
aiter.mla_reduce_v1(
logits,
attn_lse,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
max_seqlen_q,
o,
final_lse,
)
if io_transformed:
if return_logits:
logits = logits.view(-1, 1, ori_nhead, v_head_dim)
q = q.view(ori_total_s, ori_nhead, -1)
o = o.view(ori_total_s, ori_nhead, -1)
return logits, final_lse
# ---------------------------------------------------------------------------
# Inlined from aiter.ops.triton.attention.mla_decode_stage1_mxfp4
# ---------------------------------------------------------------------------
@triton.jit
def _qkt_mxfp4_kernel(
Q_fp4_ptr,
Q_scale_ptr,
K_fp4_ptr,
K_scale_ptr,
scores_ptr,
qo_indptr_ptr,
kv_indptr_ptr,
stride_q_s,
stride_q_h,
stride_q_k,
stride_q_scale_s,
stride_q_scale_k,
stride_k_kv,
stride_k_k,
stride_k_scale_kv,
stride_k_scale_k,
stride_scores_s,
stride_scores_h,
stride_scores_n,
total_s: tl.int32,
nhead: tl.int32,
qk_head_dim: tl.int32,
max_kv_len: tl.int32,
num_kv_splits: tl.int32,
sm_scale: tl.float32,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
):
batch = tl.program_id(0)
head = tl.program_id(1)
split = tl.program_id(2)
qo_start = tl.load(qo_indptr_ptr + batch)
kv_start_b = tl.load(kv_indptr_ptr + batch)
kv_end_b = tl.load(kv_indptr_ptr + batch + 1)
kv_len_b = kv_end_b - kv_start_b
kv_chunk_len = kv_len_b // num_kv_splits
kv_start = kv_start_b + split * kv_chunk_len
kv_end = kv_start + kv_chunk_len
if split == num_kv_splits - 1:
kv_end = kv_end_b
kv_chunk_len = kv_end - kv_start
if kv_chunk_len <= 0:
return
q_idx = qo_start
if q_idx >= total_s:
return
scale_k_blocks = qk_head_dim // MXFP4_BLOCK_SIZE
base_scores = (
scores_ptr
+ q_idx * stride_scores_s
+ head * stride_scores_h
+ split * max_kv_len * stride_scores_n
)
for n_start in range(0, kv_chunk_len, BLOCK_N):
n_end = tl.minimum(n_start + BLOCK_N, kv_chunk_len)
n_size = n_end - n_start
offs_n = tl.arange(0, BLOCK_N)
acc = tl.zeros((1, BLOCK_N), dtype=tl.float32)
for k_block in range(0, scale_k_blocks):
k_start = k_block * MXFP4_BLOCK_SIZE
offs_k = tl.arange(0, BLOCK_K // 2)
q_ptrs = (
Q_fp4_ptr
+ q_idx * stride_q_s
+ head * stride_q_h
+ (k_start // 2) * stride_q_k
)
q_load = tl.load(
q_ptrs,
mask=offs_k[None, :] < (qk_head_dim // 2 - k_start // 2),
other=0,
)
q_scale_ptrs = (
Q_scale_ptr
+ (q_idx * nhead + head) * stride_q_scale_s
+ k_block * stride_q_scale_k
)
q_s = tl.load(q_scale_ptrs)
q_s_bc = tl.broadcast_to(q_s, (1, 1))
k_base = (
K_fp4_ptr
+ (kv_start + n_start) * stride_k_kv
+ (k_start // 2) * stride_k_k
)
mask_n = offs_n < n_size
mask_k = offs_k < (qk_head_dim // 2 - k_start // 2)
k_load = tl.load(
k_base + offs_k[:, None] * stride_k_k + offs_n[None, :] * stride_k_kv,
mask=mask_k[:, None] & mask_n[None, :],
other=0,
)
k_scale_ptrs = (
K_scale_ptr
+ (kv_start + n_start) * stride_k_scale_kv
+ k_block * stride_k_scale_k
)
k_s = tl.load(
k_scale_ptrs + offs_n * stride_k_scale_kv, mask=mask_n, other=127
)
k_s_bc = tl.broadcast_to(k_s[None, :], (BLOCK_K // 2, BLOCK_N))
acc = tl.dot_scaled(
q_load,
q_s_bc,
"e2m1",
k_load,
k_s_bc,
"e2m1",
acc,
)
acc = acc * sm_scale
tl.store(
base_scores + n_start * stride_scores_n + offs_n * stride_scores_n,
acc[0, :],
mask=offs_n < n_size,
)
@triton.jit
def _softmax_lse_pv_f32_kernel(
scores_ptr,
V_fp32_ptr,
logits_ptr,
lse_ptr,
o_ptr,
qo_indptr_ptr,
kv_indptr_ptr,
stride_scores_s,
stride_scores_h,
stride_scores_n,
stride_v_kv,
stride_v_d,
stride_logits_s,
stride_logits_h,
stride_logits_d,
stride_lse_s,
stride_lse_h,
stride_lse_split,
stride_logits_split,
stride_o_s,
stride_o_h,
stride_o_d,
total_s: tl.int32,
nhead: tl.int32,
v_head_dim: tl.int32,
max_kv_len: tl.int32,
num_kv_splits: tl.int32,
BLOCK_D: tl.constexpr,
):
batch = tl.program_id(0)
head = tl.program_id(1)
split = tl.program_id(2)
qo_start = tl.load(qo_indptr_ptr + batch)
kv_start_b = tl.load(kv_indptr_ptr + batch)
kv_end_b = tl.load(kv_indptr_ptr + batch + 1)
kv_len_b = kv_end_b - kv_start_b
kv_chunk_len = kv_len_b // num_kv_splits
kv_start = kv_start_b + split * kv_chunk_len
kv_end = kv_start + kv_chunk_len
if split == num_kv_splits - 1:
kv_end = kv_end_b
kv_chunk_len = kv_end - kv_start
if kv_chunk_len <= 0:
return
q_idx = qo_start
if q_idx >= total_s:
return
base_scores = (
scores_ptr
+ q_idx * stride_scores_s
+ head * stride_scores_h
+ split * max_kv_len * stride_scores_n
)
m_i = -float("inf")
l_i = 0.0
BLOCK_N = 64
for n_start in range(0, kv_chunk_len, BLOCK_N):
n_end = tl.minimum(n_start + BLOCK_N, kv_chunk_len)
offs_n = tl.arange(0, BLOCK_N)
mask_n = offs_n < (n_end - n_start)
s = tl.load(
base_scores + (n_start + offs_n) * stride_scores_n,
mask=mask_n,
other=-float("inf"),
)
m_ij = tl.maximum(m_i, tl.max(s, axis=0))
p = tl.exp(s - m_ij)
p = tl.where(mask_n, p, 0.0)
alpha = tl.exp(m_i - m_ij)
l_i = l_i * alpha + tl.sum(p, axis=0)
m_i = m_ij
lse = m_i + tl.log(l_i)
acc_o = tl.zeros((BLOCK_D,), dtype=tl.float32)
offs_d = tl.arange(0, BLOCK_D)
for n_start in range(0, kv_chunk_len, BLOCK_N):
n_end = tl.minimum(n_start + BLOCK_N, kv_chunk_len)
offs_n = tl.arange(0, BLOCK_N)
mask_n = offs_n < (n_end - n_start)
s = tl.load(
base_scores + (n_start + offs_n) * stride_scores_n,
mask=mask_n,
other=0.0,
)
p = tl.exp(s - m_i) / l_i
p = tl.where(mask_n, p, 0.0)
for ni in range(BLOCK_N):
if n_start + ni >= kv_chunk_len:
break
p_val = tl.load(base_scores + (n_start + ni) * stride_scores_n)
p_val = tl.exp(p_val - m_i) / l_i
v_row = tl.load(
V_fp32_ptr
+ (kv_start + n_start + ni) * stride_v_kv
+ offs_d * stride_v_d,
mask=offs_d < v_head_dim,
other=0.0,
)
acc_o += p_val * v_row
lse_off = q_idx * stride_lse_s + split * stride_lse_split + head * stride_lse_h
tl.store(lse_ptr + lse_off, lse)
logits_off = (
q_idx * stride_logits_s + split * stride_logits_split + head * stride_logits_h
)
tl.store(
logits_ptr + logits_off + offs_d * stride_logits_d,
acc_o,
mask=offs_d < v_head_dim,
)
if num_kv_splits == 1:
tl.store(
o_ptr + q_idx * stride_o_s + head * stride_o_h + offs_d * stride_o_d,
acc_o,
mask=offs_d < v_head_dim,
)
def _dequant_mxfp4_to_f32(
x_fp4: torch.Tensor, scale_e8m0: torch.Tensor, dim: int
) -> torch.Tensor:
from aiter.utility import fp4_utils
vals = fp4_utils.mxfp4_to_f32(x_fp4)
scale = scale_e8m0.view(torch.uint8).to(torch.float32)
scale = torch.pow(2.0, 127.0 - scale)
if scale.dim() == 2:
scale = scale[:, : dim // MXFP4_BLOCK_SIZE].repeat_interleave(
MXFP4_BLOCK_SIZE, dim=1
)
return vals * scale
def _gather_kv_from_paged(
kv_buffer: torch.Tensor,
kv_indices: torch.Tensor,
page_size: int,
nhead_kv: int,
qk_head_dim: int,
v_head_dim: int,
) -> tuple[torch.Tensor, torch.Tensor]:
total_kv = kv_indices.shape[0]
if kv_buffer.dim() == 4:
num_page, ps, nk, feat = kv_buffer.shape
kv_flat = kv_buffer.reshape(num_page * ps, nk, feat)
flat_idx = (
kv_indices * page_size
+ torch.arange(total_kv, device=kv_indices.device, dtype=kv_indices.dtype)
% page_size
)
kv_gathered = kv_flat[flat_idx]
else:
kv_flat = kv_buffer
kv_gathered = kv_flat[kv_indices]
if kv_gathered.shape[-1] >= qk_head_dim + v_head_dim:
K = kv_gathered[..., :qk_head_dim].contiguous()
V = kv_gathered[..., qk_head_dim : qk_head_dim + v_head_dim].contiguous()
else:
K = kv_gathered.contiguous()
V = kv_gathered[..., :v_head_dim].contiguous()
return K, V
def mla_decode_stage1_mxfp4(
q: torch.Tensor,
kv_buffer: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
kv_indices: torch.Tensor,
kv_last_page_lens: torch.Tensor,
num_kv_splits_indptr: torch.Tensor,
max_seqlen_q: int,
page_size: int,
nhead_kv: int,
sm_scale: float,
logits: torch.Tensor,
attn_lse: torch.Tensor,
o: torch.Tensor,
) -> None:
assert arch_info.is_fp4_avail(), "MXFP4 requires gfx950"
device = q.device
total_s, nhead, qk_head_dim = q.shape
v_head_dim = o.shape[-1]
bs = qo_indptr.shape[0] - 1
total_kv = kv_indices.shape[0]
num_kv_splits = (num_kv_splits_indptr[1] - num_kv_splits_indptr[0]).item()
assert qk_head_dim % MXFP4_BLOCK_SIZE == 0
assert v_head_dim % MXFP4_BLOCK_SIZE == 0
q_flat = q.reshape(-1, qk_head_dim).to(torch.bfloat16)
Q_fp4, Q_scale = dynamic_mxfp4_quant(q_flat)
Q_fp4 = Q_fp4.reshape(total_s, nhead, -1)
scale_n_q = qk_head_dim // MXFP4_BLOCK_SIZE
Q_scale = Q_scale.reshape(-1, scale_n_q)[: total_s * nhead].contiguous()
K_gather, V_gather = _gather_kv_from_paged(
kv_buffer, kv_indices, page_size, nhead_kv, qk_head_dim, v_head_dim
)
K_flat = K_gather.reshape(-1, qk_head_dim).to(torch.bfloat16)
V_flat = V_gather.reshape(-1, v_head_dim).to(torch.bfloat16)
K_fp4, K_scale = dynamic_mxfp4_quant(K_flat)
V_fp4, V_scale = dynamic_mxfp4_quant(V_flat)
V_fp32 = _dequant_mxfp4_to_f32(V_fp4, V_scale, v_head_dim)
max_kv_len = (kv_indptr[1:] - kv_indptr[:-1]).max().item()
max_kv_chunk = (max_kv_len + num_kv_splits - 1) // num_kv_splits
scores = torch.empty(
(total_s, nhead, num_kv_splits * max_kv_chunk),
dtype=torch.float32,
device=device,
)
scores.fill_(-1e9)
BLOCK_N = 64
BLOCK_K = 32
grid_qkt = (bs, nhead, num_kv_splits)
_qkt_mxfp4_kernel[grid_qkt](
Q_fp4,
Q_scale,
K_fp4,
K_scale,
scores,
qo_indptr,
kv_indptr,
stride_q_s=Q_fp4.stride(0),
stride_q_h=Q_fp4.stride(1),
stride_q_k=Q_fp4.stride(2),
stride_q_scale_s=Q_scale.stride(0),
stride_q_scale_k=Q_scale.stride(1),
stride_k_kv=K_fp4.stride(0),
stride_k_k=K_fp4.stride(1),
stride_k_scale_kv=K_scale.stride(0),
stride_k_scale_k=K_scale.stride(1),
stride_scores_s=scores.stride(0),
stride_scores_h=scores.stride(1),
stride_scores_n=scores.stride(2),
total_s=total_s,
nhead=nhead,
qk_head_dim=qk_head_dim,
max_kv_len=max_kv_chunk,
num_kv_splits=num_kv_splits,
sm_scale=sm_scale,
BLOCK_N=BLOCK_N,
BLOCK_K=BLOCK_K,
num_warps=4,
)
BLOCK_D = triton.next_power_of_2(v_head_dim)
_softmax_lse_pv_f32_kernel[grid_qkt](
scores,
V_fp32,
logits,
attn_lse,
o,
qo_indptr,
kv_indptr,
stride_scores_s=scores.stride(0),
stride_scores_h=scores.stride(1),
stride_scores_n=scores.stride(2),
stride_v_kv=V_fp32.stride(0),
stride_v_d=V_fp32.stride(1),
stride_logits_s=logits.stride(0),
stride_logits_h=logits.stride(2),
stride_logits_d=logits.stride(3),
stride_lse_s=attn_lse.stride(0),
stride_lse_h=attn_lse.stride(2),
stride_lse_split=attn_lse.stride(1),
stride_logits_split=logits.stride(1),
stride_o_s=o.stride(0),
stride_o_h=o.stride(1),
stride_o_d=o.stride(2),
total_s=total_s,
nhead=nhead,
v_head_dim=v_head_dim,
max_kv_len=max_kv_chunk,
num_kv_splits=num_kv_splits,
BLOCK_D=BLOCK_D,
num_warps=4,
)
# ---------------------------------------------------------------------------
# Persistent mode metadata and wrapper (calls local mla_decode_fwd)
# ---------------------------------------------------------------------------
def _make_mla_decode_metadata(
batch_size: int,
max_q_len: int,
nhead: int,
nhead_kv: int,
q_dtype: torch.dtype,
kv_dtype: torch.dtype,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
kv_last_page_len: torch.Tensor,
num_kv_splits: int = NUM_KV_SPLITS,
):
"""Allocate and populate work buffers for persistent mla_decode_fwd."""
info = get_mla_metadata_info_v1(
batch_size,
max_q_len,
nhead,
q_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,
kv_last_page_len,
nhead // nhead_kv,
nhead_kv,
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=True,
dtype_q=q_dtype,
dtype_kv=kv_dtype,
)
return {
"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,
}
def _aiter_mla_decode(
q: torch.Tensor,
kv_buffer: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
config: dict,
q_scale: torch.Tensor | None = None,
kv_scale: torch.Tensor | None = None,
) -> torch.Tensor:
batch_size = config["batch_size"]
nq = config["num_heads"]
nkv = config["num_kv_heads"]
dq = config["qk_head_dim"]
dv = config["v_head_dim"]
q_seq_len = config["q_seq_len"]
total_kv_len = int(kv_indptr[-1].item())
kv_indices = torch.arange(total_kv_len, dtype=torch.int32, device="cuda")
kv_buffer_4d = kv_buffer.view(
kv_buffer.shape[0], PAGE_SIZE, nkv, kv_buffer.shape[-1]
)
max_q_len = q_seq_len
kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
meta = _make_mla_decode_metadata(
batch_size,
max_q_len,
nq,
nkv,
q.dtype,
kv_buffer.dtype,
qo_indptr,
kv_indptr,
kv_last_page_len,
num_kv_splits=NUM_KV_SPLITS,
)
o = torch.empty((q.shape[0], nq, dv), dtype=torch.bfloat16, device="cuda")
mla_decode_fwd(
q.view(-1, nq, dq),
kv_buffer_4d,
o,
qo_indptr,
kv_indptr,
kv_indices,
kv_last_page_len,
max_q_len,
page_size=PAGE_SIZE,
nhead_kv=nkv,
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,
**meta,
)
return o
# ---------------------------------------------------------------------------
# Strategy A: mxfp4 KV (dequant then attend); fp8 / bf16 strategies
# ---------------------------------------------------------------------------
def _mla_decode_strategy_a(
q: torch.Tensor,
kv_buffer_mxfp4: torch.Tensor,
kv_scale_mxfp4: torch.Tensor,
qo_indptr: torch.Tensor,
kv_indptr: torch.Tensor,
config: dict,
) -> torch.Tensor:
kv_orig_shape = (kv_buffer_mxfp4.shape[0], kv_buffer_mxfp4.shape[1], QK_HEAD_DIM)
kv_bf16 = dequantize_mxfp4(kv_buffer_mxfp4, kv_scale_mxfp4, kv_orig_shape)
return _aiter_mla_decode(
q,
kv_bf16,
qo_indptr,
kv_indptr,
config,
q_scale=None,
kv_scale=None,
)
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
if QKV_DTYPE == "mxfp4":
kv_buffer_mxfp4, kv_scale_mxfp4 = kv_data["mxfp4"]
return _mla_decode_strategy_a(
q,
kv_buffer_mxfp4,
kv_scale_mxfp4,
qo_indptr,
kv_indptr,
config,
)
elif QKV_DTYPE == "fp8":
q_input, q_scale = quantize_fp8(q)
kv_buffer_fp8, kv_scale = kv_data["fp8"]
return _aiter_mla_decode(
q_input,
kv_buffer_fp8,
qo_indptr,
kv_indptr,
config,
q_scale=q_scale,
kv_scale=kv_scale,
)
elif QKV_DTYPE == "bf16":
q_input, q_scale = q, None
kv_input, kv_scale = kv_data["bf16"], None
return _aiter_mla_decode(
q_input,
kv_input,
qo_indptr,
kv_indptr,
config,
q_scale=q_scale,
kv_scale=kv_scale,
)
else:
raise ValueError(f"Invalid QKV_DTYPE: {QKV_DTYPE}")
scrolls · 1119 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