submission 746469
chenxingqiang · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 590 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-746469?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:f8759a3cfaf68d074b30374cf698076912daf4ba7589f1088ae9c3de735a9a01
license declaredunknown
license concludedunknown
authorschenxingqiang
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
s = tl.dot(q_lora, tl.trans(k_lora)) + tl.dot(q_rope, tl.trans(k_rope))num-warps = 4
num_warps=4, num_stages=2,online-softmax
m_new = tl.maximum(m_i, tl.max(s, axis=1))stages = 2
num_warps=4, num_stages=2,Kernel source
submission.py590 lines
"""
Hybrid MLA decode: custom Triton for short KV, ASM for long KV.
"""
import os
import torch
import triton
import triton.language as tl
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.mla import mla_decode_fwd
from aiter.ops.triton.quant import dynamic_per_tensor_quant_fp8_i8
from aiter.ops.attention import mla_decode_stage1_asm_fwd, mla_reduce_v1
from aiter.utility.fp4_utils import mxfp4_to_f32, e8m0_to_f32
KV_LORA_RANK = 512
QK_ROPE_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_DIM
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
_CACHE: dict = {}
_ASM_CACHE: dict = {}
_MXFP4_CACHE: dict = {}
_CACHE_LIMIT = 32
@triton.jit
def _decode_mxfp4_nibbles(x):
mag = x & 0x7
sign = x & 0x8
out = tl.where(mag == 0, 0.0, 0.0)
out = tl.where(mag == 1, 0.5, out)
out = tl.where(mag == 2, 1.0, out)
out = tl.where(mag == 3, 1.5, out)
out = tl.where(mag == 4, 2.0, out)
out = tl.where(mag == 5, 3.0, out)
out = tl.where(mag == 6, 4.0, out)
out = tl.where(mag == 7, 6.0, out)
return tl.where(sign != 0, -out, out)
@triton.jit
def _mla_decode_attn(
Q, KV, Out,
Partial_O, Partial_M, Partial_L,
qo_indptr, kv_indptr,
kv_scale_ptr,
sm_scale: tl.constexpr,
stride_q_s, stride_q_h,
stride_kv_s,
stride_o_s, stride_o_h,
NHEADS: tl.constexpr,
D_LORA: tl.constexpr,
D_ROPE: tl.constexpr,
BLOCK_KV: tl.constexpr,
NUM_SPLITS: tl.constexpr,
):
pid = tl.program_id(0)
pid_b = pid // NUM_SPLITS
pid_s = pid % NUM_SPLITS
qo_start = tl.load(qo_indptr + pid_b)
kv_start = tl.load(kv_indptr + pid_b)
kv_end = tl.load(kv_indptr + pid_b + 1)
kv_len = kv_end - kv_start
chunk = tl.cdiv(kv_len, NUM_SPLITS)
my_start = kv_start + pid_s * chunk
my_end = tl.minimum(kv_start + (pid_s + 1) * chunk, kv_end)
my_len = my_end - my_start
kv_scale = tl.load(kv_scale_ptr).to(tl.float32)
qk_combined = kv_scale * sm_scale
offs_h = tl.arange(0, NHEADS)
offs_lora = tl.arange(0, D_LORA)
offs_rope = tl.arange(0, D_ROPE)
q_lora = tl.load(
Q + qo_start * stride_q_s + offs_h[:, None] * stride_q_h + offs_lora[None, :]
)
q_rope = tl.load(
Q + qo_start * stride_q_s + offs_h[:, None] * stride_q_h + D_LORA + offs_rope[None, :]
)
m_i = tl.full([NHEADS], value=float("-inf"), dtype=tl.float32)
l_i = tl.zeros([NHEADS], dtype=tl.float32)
acc = tl.zeros([NHEADS, D_LORA], dtype=tl.float32)
offs_kv = tl.arange(0, BLOCK_KV)
for kv_off in range(0, my_len, BLOCK_KV):
valid = (kv_off + offs_kv) < my_len
base = my_start + kv_off
k_lora = tl.load(
KV + (base + offs_kv[:, None]) * stride_kv_s + offs_lora[None, :],
mask=valid[:, None], other=0.0,
).to(tl.bfloat16)
k_rope = tl.load(
KV + (base + offs_kv[:, None]) * stride_kv_s + D_LORA + offs_rope[None, :],
mask=valid[:, None], other=0.0,
).to(tl.bfloat16)
s = tl.dot(q_lora, tl.trans(k_lora)) + tl.dot(q_rope, tl.trans(k_rope))
s *= qk_combined
s = tl.where(valid[None, :], s, float("-inf"))
m_new = tl.maximum(m_i, tl.max(s, axis=1))
alpha = tl.exp(m_i - m_new)
p = tl.exp(s - m_new[:, None])
acc = acc * alpha[:, None] + tl.dot(p.to(tl.bfloat16), k_lora)
l_i = l_i * alpha + tl.sum(p, axis=1)
m_i = m_new
if NUM_SPLITS == 1:
acc = acc * (kv_scale / l_i[:, None])
tl.store(
Out + qo_start * stride_o_s + offs_h[:, None] * stride_o_h + offs_lora[None, :],
acc.to(tl.bfloat16),
)
else:
part_base = pid * NHEADS
tl.store(
Partial_O + pid * NHEADS * D_LORA + offs_h[:, None] * D_LORA + offs_lora[None, :],
acc,
)
tl.store(Partial_M + part_base + offs_h, m_i)
tl.store(Partial_L + part_base + offs_h, l_i)
@triton.jit
def _mla_decode_attn_mxfp4(
Q, KV, KV_SCALES, Out,
Partial_O, Partial_M, Partial_L,
qo_indptr, kv_indptr,
sm_scale: tl.constexpr,
stride_q_s, stride_q_h,
stride_kv_s, stride_scale_s,
stride_o_s, stride_o_h,
NHEADS: tl.constexpr,
D_LORA: tl.constexpr,
D_ROPE: tl.constexpr,
BLOCK_KV: tl.constexpr,
NUM_SPLITS: tl.constexpr,
):
pid = tl.program_id(0)
pid_b = pid // NUM_SPLITS
pid_s = pid % NUM_SPLITS
qo_start = tl.load(qo_indptr + pid_b)
kv_start = tl.load(kv_indptr + pid_b)
kv_end = tl.load(kv_indptr + pid_b + 1)
kv_len = kv_end - kv_start
chunk = tl.cdiv(kv_len, NUM_SPLITS)
my_start = kv_start + pid_s * chunk
my_end = tl.minimum(kv_start + (pid_s + 1) * chunk, kv_end)
my_len = my_end - my_start
offs_h = tl.arange(0, NHEADS)
offs_lora = tl.arange(0, D_LORA)
offs_rope = tl.arange(0, D_ROPE)
q_lora = tl.load(
Q + qo_start * stride_q_s + offs_h[:, None] * stride_q_h + offs_lora[None, :]
)
q_rope = tl.load(
Q + qo_start * stride_q_s + offs_h[:, None] * stride_q_h + D_LORA + offs_rope[None, :]
)
m_i = tl.full([NHEADS], value=float("-inf"), dtype=tl.float32)
l_i = tl.zeros([NHEADS], dtype=tl.float32)
acc = tl.zeros([NHEADS, D_LORA], dtype=tl.float32)
offs_kv = tl.arange(0, BLOCK_KV)
packed_lora_idx = offs_lora // 2
packed_rope_idx = (D_LORA + offs_rope) // 2
scale_lora_idx = offs_lora // 32
scale_rope_idx = (D_LORA + offs_rope) // 32
lora_is_low = (offs_lora % 2) == 0
rope_is_low = ((D_LORA + offs_rope) % 2) == 0
for kv_off in range(0, my_len, BLOCK_KV):
valid = (kv_off + offs_kv) < my_len
base = my_start + kv_off
packed_lora = tl.load(
KV + (base + offs_kv[:, None]) * stride_kv_s + packed_lora_idx[None, :],
mask=valid[:, None], other=0,
)
packed_rope = tl.load(
KV + (base + offs_kv[:, None]) * stride_kv_s + packed_rope_idx[None, :],
mask=valid[:, None], other=0,
)
nib_lora = tl.where(lora_is_low[None, :], packed_lora & 0xF, packed_lora >> 4)
nib_rope = tl.where(rope_is_low[None, :], packed_rope & 0xF, packed_rope >> 4)
scales_lora = tl.load(
KV_SCALES + (base + offs_kv[:, None]) * stride_scale_s + scale_lora_idx[None, :],
mask=valid[:, None], other=127,
)
scales_rope = tl.load(
KV_SCALES + (base + offs_kv[:, None]) * stride_scale_s + scale_rope_idx[None, :],
mask=valid[:, None], other=127,
)
k_lora = _decode_mxfp4_nibbles(nib_lora.to(tl.uint8)) * tl.exp2(scales_lora.to(tl.float32) - 127.0)
k_rope = _decode_mxfp4_nibbles(nib_rope.to(tl.uint8)) * tl.exp2(scales_rope.to(tl.float32) - 127.0)
s = tl.dot(q_lora, tl.trans(k_lora.to(tl.bfloat16))) + tl.dot(q_rope, tl.trans(k_rope.to(tl.bfloat16)))
s *= sm_scale
s = tl.where(valid[None, :], s, float("-inf"))
m_new = tl.maximum(m_i, tl.max(s, axis=1))
alpha = tl.exp(m_i - m_new)
p = tl.exp(s - m_new[:, None])
acc = acc * alpha[:, None] + tl.dot(p.to(tl.bfloat16), k_lora.to(tl.bfloat16))
l_i = l_i * alpha + tl.sum(p, axis=1)
m_i = m_new
if NUM_SPLITS == 1:
acc = acc / l_i[:, None]
tl.store(
Out + qo_start * stride_o_s + offs_h[:, None] * stride_o_h + offs_lora[None, :],
acc.to(tl.bfloat16),
)
else:
part_base = pid * NHEADS
tl.store(
Partial_O + pid * NHEADS * D_LORA + offs_h[:, None] * D_LORA + offs_lora[None, :],
acc,
)
tl.store(Partial_M + part_base + offs_h, m_i)
tl.store(Partial_L + part_base + offs_h, l_i)
@triton.jit
def _mla_decode_reduce(
Partial_O, Partial_M, Partial_L,
Out, qo_indptr, kv_scale_ptr,
stride_o_s, stride_o_h,
NHEADS: tl.constexpr,
D_LORA: tl.constexpr,
NUM_SPLITS: tl.constexpr,
BLOCK_V: tl.constexpr,
):
pid_b = tl.program_id(0)
pid_h = tl.program_id(1)
qo_start = tl.load(qo_indptr + pid_b)
kv_scale = tl.load(kv_scale_ptr).to(tl.float32)
offs_v = tl.arange(0, BLOCK_V)
mask_v = offs_v < D_LORA
m_global = float("-inf")
l_global = 0.0
acc = tl.zeros([BLOCK_V], dtype=tl.float32)
for s in range(NUM_SPLITS):
part_idx = pid_b * NUM_SPLITS + s
m_s = tl.load(Partial_M + part_idx * NHEADS + pid_h)
l_s = tl.load(Partial_L + part_idx * NHEADS + pid_h)
po = tl.load(
Partial_O + part_idx * NHEADS * D_LORA + pid_h * D_LORA + offs_v,
mask=mask_v, other=0.0,
)
m_new = tl.maximum(m_global, m_s)
alpha = tl.exp(m_global - m_new)
beta = tl.exp(m_s - m_new)
acc = acc * alpha + po * beta
l_global = l_global * alpha + l_s * beta
m_global = m_new
acc = acc * (kv_scale / l_global)
tl.store(
Out + qo_start * stride_o_s + pid_h * stride_o_h + offs_v,
acc.to(tl.bfloat16),
mask=mask_v,
)
def _get_num_splits(batch_size, kv_seq_len):
target_blocks = 256
blocks_per_batch = max(1, target_blocks // batch_size)
max_splits = max(1, kv_seq_len // 64)
return min(blocks_per_batch, max_splits)
def _select_asm_splits(batch_size, kv_seq_len):
if kv_seq_len >= 8192:
if batch_size >= 64:
return 64
return 32 if batch_size >= 16 else 16
if kv_seq_len >= 4096:
return 32
return 16
def _select_mxfp4_splits(batch_size, kv_seq_len):
if kv_seq_len >= 8192:
if batch_size >= 64:
return 32
return 16
if kv_seq_len >= 4096:
return 16
return 8
def _build_asm_cache(batch_size, kv_seq_len, num_kv_splits, nq, nkv, dq, dv,
q_dtype, kv_dtype, qo_indptr, kv_indptr, device):
total_kv = batch_size * kv_seq_len
kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)
kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
kv_gran = 128 if kv_seq_len >= 8192 else (64 if kv_seq_len >= 4096 else 16)
info = get_mla_metadata_info_v1(
batch_size, 1, nq, 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=device) for s, t in info]
wm, wi, wis, ri, rfm, rpm = work
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_last_page_len,
nq // nkv, nkv, True, wm, wis, wi, ri, rfm, rpm,
page_size=PAGE_SIZE, kv_granularity=kv_gran,
max_seqlen_qo=1, uni_seqlen_qo=1,
fast_mode=False, max_split_per_batch=num_kv_splits,
intra_batch_mode=True, dtype_q=q_dtype, dtype_kv=kv_dtype,
)
n_partial = rpm.size(0)
return {
"kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,
"wm": wm, "wi": wi, "wis": wis, "ri": ri, "rfm": rfm, "rpm": rpm,
"logits": torch.empty((n_partial, 1, nq, dv), dtype=torch.float32, device=device),
"attn_lse": torch.empty((n_partial, 1, nq, 1), dtype=torch.float32, device=device),
"q_fp8": torch.empty((batch_size * nq, dq), dtype=FP8_DTYPE, device=device),
"q_scale": torch.empty((1,), dtype=torch.float32, device=device),
"o": torch.empty((batch_size, nq, dv), dtype=torch.bfloat16, device=device),
}
def _run_triton_decode(q, kv_buffer_fp8, kv_scale, qo_indptr, kv_indptr,
batch_size, kv_seq_len, nq, dv, block_kv, num_splits, cache_tag):
cache_key = (cache_tag, batch_size, kv_seq_len, num_splits, block_kv)
c = _CACHE.get(cache_key)
if c is None:
o = torch.empty((batch_size, nq, dv), dtype=torch.bfloat16, device=q.device)
if num_splits > 1:
tp = batch_size * num_splits
po = torch.empty((tp, nq, KV_LORA_RANK), dtype=torch.float32, device=q.device)
pm = torch.empty((tp, nq), dtype=torch.float32, device=q.device)
pl = torch.empty((tp, nq), dtype=torch.float32, device=q.device)
else:
po = pm = pl = None
c = {"o": o, "po": po, "pm": pm, "pl": pl}
if len(_CACHE) >= _CACHE_LIMIT:
_CACHE.clear()
_CACHE[cache_key] = c
kv_flat = kv_buffer_fp8.view(-1, kv_buffer_fp8.shape[-1])
_mla_decode_attn[(batch_size * num_splits,)](
q, kv_flat, c["o"], c["po"], c["pm"], c["pl"],
qo_indptr, kv_indptr, kv_scale, SM_SCALE,
q.stride(0), q.stride(1), kv_flat.stride(0),
c["o"].stride(0), c["o"].stride(1),
NHEADS=nq, D_LORA=KV_LORA_RANK, D_ROPE=QK_ROPE_DIM,
BLOCK_KV=block_kv, NUM_SPLITS=num_splits,
num_warps=4, num_stages=2,
)
if num_splits > 1:
_mla_decode_reduce[(batch_size, nq)](
c["po"], c["pm"], c["pl"], c["o"],
qo_indptr, kv_scale,
c["o"].stride(0), c["o"].stride(1),
NHEADS=nq, D_LORA=KV_LORA_RANK,
NUM_SPLITS=num_splits, BLOCK_V=KV_LORA_RANK,
num_warps=4,
)
return c["o"]
_FWD_CACHE: dict = {}
def _run_mla_decode_fwd(q, kv_buffer_fp8, kv_scale, qo_indptr, kv_indptr,
batch_size, kv_seq_len, nq, nkv, dq, dv):
"""Use high-level mla_decode_fwd (same as reference) for best perf."""
num_kv_splits = 32 if kv_seq_len >= 4096 else 16
cache_key = ("fwd", batch_size, kv_seq_len, num_kv_splits)
c = _FWD_CACHE.get(cache_key)
device = q.device
if c is None:
total_kv = batch_size * kv_seq_len
kv_indices = torch.arange(total_kv, dtype=torch.int32, device=device)
kv_last_page_len = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
info = get_mla_metadata_info_v1(
batch_size, 1, nq, FP8_DTYPE, kv_buffer_fp8.dtype,
is_sparse=False, fast_mode=False,
num_kv_splits=num_kv_splits, intra_batch_mode=True,
)
work = [torch.empty(s, dtype=t, device=device) for s, t in info]
wm, wi, wis, ri, rfm, rpm = work
get_mla_metadata_v1(
qo_indptr, kv_indptr, kv_last_page_len,
nq // nkv, nkv, True, wm, wis, wi, ri, rfm, rpm,
page_size=PAGE_SIZE, kv_granularity=16,
max_seqlen_qo=1, uni_seqlen_qo=1,
fast_mode=False, max_split_per_batch=num_kv_splits,
intra_batch_mode=True, dtype_q=FP8_DTYPE, dtype_kv=kv_buffer_fp8.dtype,
)
c = {
"kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,
"wm": wm, "wi": wi, "wis": wis, "ri": ri, "rfm": rfm, "rpm": rpm,
"q_fp8": torch.empty((batch_size * nq, dq), dtype=FP8_DTYPE, device=device),
"q_scale": torch.empty((1,), dtype=torch.float32, device=device),
"o": torch.empty((batch_size, nq, dv), dtype=torch.bfloat16, device=device),
}
if len(_FWD_CACHE) >= _CACHE_LIMIT:
_FWD_CACHE.clear()
_FWD_CACHE[cache_key] = c
q_2d = q.view(-1, dq)
if not q_2d.is_contiguous():
q_2d = q_2d.contiguous()
dynamic_per_tensor_quant_fp8_i8(c["q_fp8"], q_2d, c["q_scale"])
kv_4d = kv_buffer_fp8.view(kv_buffer_fp8.shape[0], PAGE_SIZE, nkv, kv_buffer_fp8.shape[-1])
mla_decode_fwd(
c["q_fp8"].view(batch_size, nq, dq),
kv_4d,
c["o"],
qo_indptr, kv_indptr,
c["kv_indices"], c["kv_last_page_len"],
1, page_size=PAGE_SIZE, nhead_kv=nkv, sm_scale=SM_SCALE,
logit_cap=0.0, num_kv_splits=num_kv_splits,
q_scale=c["q_scale"], kv_scale=kv_scale,
intra_batch_mode=True,
work_meta_data=c["wm"], work_indptr=c["wi"], work_info_set=c["wis"],
reduce_indptr=c["ri"], reduce_final_map=c["rfm"], reduce_partial_map=c["rpm"],
)
return c["o"]
def _run_asm_decode(q, kv_buffer_fp8, kv_scale, qo_indptr, kv_indptr,
batch_size, kv_seq_len, nq, nkv, dq, dv):
num_kv_splits = _select_asm_splits(batch_size, kv_seq_len)
cache_key = ("asm", batch_size, kv_seq_len, num_kv_splits)
c = _ASM_CACHE.get(cache_key)
if c is None:
c = _build_asm_cache(
batch_size, kv_seq_len, num_kv_splits, nq, nkv, dq, dv,
FP8_DTYPE, kv_buffer_fp8.dtype,
qo_indptr, kv_indptr, q.device,
)
if len(_ASM_CACHE) >= _CACHE_LIMIT:
_ASM_CACHE.clear()
_ASM_CACHE[cache_key] = c
q_2d = q.view(-1, dq)
if not q_2d.is_contiguous():
q_2d = q_2d.contiguous()
dynamic_per_tensor_quant_fp8_i8(c["q_fp8"], q_2d, c["q_scale"])
kv_4d = kv_buffer_fp8.view(kv_buffer_fp8.shape[0], PAGE_SIZE, nkv, kv_buffer_fp8.shape[-1])
mla_decode_stage1_asm_fwd(
c["q_fp8"].view(batch_size, nq, dq), kv_4d,
qo_indptr, kv_indptr,
c["kv_indices"], c["kv_last_page_len"],
None, c["wm"], c["wi"], c["wis"],
1, PAGE_SIZE, nkv, SM_SCALE,
c["logits"], c["attn_lse"], c["o"],
c["q_scale"], kv_scale,
)
mla_reduce_v1(
c["logits"], c["attn_lse"],
c["ri"], c["rfm"], c["rpm"],
1, c["o"], None,
)
return c["o"]
def _dequantize_mxfp4_kv(kv_buffer_mxfp4, kv_scale_mxfp4):
total_kv = kv_buffer_mxfp4.shape[0]
key = (
kv_buffer_mxfp4.data_ptr(),
kv_scale_mxfp4.data_ptr(),
total_kv,
kv_buffer_mxfp4.device,
)
cached = _MXFP4_CACHE.get(key)
if cached is not None:
return cached
num_blocks = QK_HEAD_DIM // 32
kv_fp32 = mxfp4_to_f32(kv_buffer_mxfp4.view(total_kv, QK_HEAD_DIM // 2))
scale_f32 = e8m0_to_f32(kv_scale_mxfp4)[:total_kv, :num_blocks]
kv_fp32 = kv_fp32.view(total_kv, num_blocks, 32) * scale_f32.unsqueeze(-1)
kv_bf16 = kv_fp32.view(total_kv, 1, QK_HEAD_DIM).to(torch.bfloat16)
if len(_MXFP4_CACHE) >= 4:
_MXFP4_CACHE.clear()
_MXFP4_CACHE[key] = kv_bf16
return kv_bf16
def _run_triton_decode_mxfp4(q, kv_buffer_mxfp4, kv_scale_mxfp4, qo_indptr, kv_indptr,
batch_size, kv_seq_len, nq, dv, num_splits, cache_tag):
cache_key = (cache_tag, batch_size, kv_seq_len, num_splits)
c = _CACHE.get(cache_key)
if c is None:
o = torch.empty((batch_size, nq, dv), dtype=torch.bfloat16, device=q.device)
if num_splits > 1:
tp = batch_size * num_splits
po = torch.empty((tp, nq, KV_LORA_RANK), dtype=torch.float32, device=q.device)
pm = torch.empty((tp, nq), dtype=torch.float32, device=q.device)
pl = torch.empty((tp, nq), dtype=torch.float32, device=q.device)
else:
po = pm = pl = None
c = {"o": o, "po": po, "pm": pm, "pl": pl}
if len(_CACHE) >= _CACHE_LIMIT:
_CACHE.clear()
_CACHE[cache_key] = c
# Cast to uint8 for Triton (FP4 dtypes not supported in pointer canonicalization)
kv_flat = kv_buffer_mxfp4.view(torch.uint8).view(-1, kv_buffer_mxfp4.shape[-1])
kv_scales = kv_scale_mxfp4.view(torch.uint8).view(kv_scale_mxfp4.shape[0], -1)
_mla_decode_attn_mxfp4[(batch_size * num_splits,)](
q, kv_flat, kv_scales, c["o"], c["po"], c["pm"], c["pl"],
qo_indptr, kv_indptr, SM_SCALE,
q.stride(0), q.stride(1),
kv_flat.stride(0), kv_scales.stride(0),
c["o"].stride(0), c["o"].stride(1),
NHEADS=nq, D_LORA=KV_LORA_RANK, D_ROPE=QK_ROPE_DIM,
BLOCK_KV=128, NUM_SPLITS=num_splits,
num_warps=4, num_stages=2,
)
if num_splits > 1:
unit_scale = torch.ones((1,), dtype=torch.float32, device=q.device)
_mla_decode_reduce[(batch_size, nq)](
c["po"], c["pm"], c["pl"], c["o"],
qo_indptr, unit_scale,
c["o"].stride(0), c["o"].stride(1),
NHEADS=nq, D_LORA=KV_LORA_RANK,
NUM_SPLITS=num_splits, BLOCK_V=KV_LORA_RANK,
num_warps=4,
)
return c["o"]
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
nq = config["num_heads"]
nkv = config["num_kv_heads"]
dq = config["qk_head_dim"]
dv = config["v_head_dim"]
kv_seq_len = config["kv_seq_len"]
kv_buffer_fp8, kv_scale = kv_data["fp8"]
# Short KV: Use custom Triton with FP8 KV
if kv_seq_len <= 2048:
num_splits = _get_num_splits(batch_size, kv_seq_len)
return _run_triton_decode(
q, kv_buffer_fp8, kv_scale, qo_indptr, kv_indptr,
batch_size, kv_seq_len, nq, dv,
block_kv=64, num_splits=num_splits, cache_tag="triton-short",
)
# Long KV: ASM decode path
return _run_asm_decode(
q, kv_buffer_fp8, kv_scale, qo_indptr, kv_indptr,
batch_size, kv_seq_len, nq, nkv, dq, dv,
)
scrolls · 590 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