submission 668813
jiajia931 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 838 lines, June 9 Researcher Reciprocity License v1.0.
submission_v0092a.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-668813?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:3677aa4e83aa0c9793dca7c69d1cd55f4ad80de05dbffcd73ec5f8ecf45aab0a
license declaredunknown
license concludedunknown
authorsjiajia931
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
v0092a: three-way hybrid a8w8 + a16w8 + mxfp4.mma
acc_e0 += tl.dot(p_f16, vlo0_s, out_dtype=tl.float32)num-warps = 4
num_warps=4,online-softmax
m_new = tl.maximum(m_i, qk_max)persistent-kernel
None, # num_kv_splits_indptr = None -> persistent modestages = 1
num_stages=1,tile-n = 32
BLOCK_N = 32Kernel source
submission_v0092a.py838 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
# submission_version: v0092a
"""
v0092a: three-way hybrid a8w8 + a16w8 + mxfp4.
Based on v0092 mxfp4 kernel (KV reuse, log2 domain, scale register gather).
Strategy per shape:
- bs=4: a8w8 (tiny Q quant cost, fp8 Q saves read bandwidth)
- bs>=32 kv=1024: a16w8 (skip Q quant, dominates on short sequences)
- bs>=32 kv=8192: mxfp4 (hardware dot_scaled, 2x KV bandwidth savings)
"""
import os
os.environ.setdefault("TRITON_CACHE_DIR", "/tmp/triton_cache_v0092a")
from typing import Any
import torch
import triton
import triton.language as tl
import aiter as aiter_module
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
try:
from aiter.utility.fp4_utils import dynamic_mxfp4_quant
HAS_MXFP4 = True
except ImportError:
HAS_MXFP4 = False
try:
from aiter.ops.quant import dynamic_per_tensor_quant
HAS_FUSED_QUANT = True
except ImportError:
HAS_FUSED_QUANT = False
input_t = Any
output_t = Any
NUM_HEADS = 16
NUM_KV_HEADS = 1
KV_LORA_RANK = 512
QK_ROPE_HEAD_DIM = 64
QK_HEAD_DIM = KV_LORA_RANK + QK_ROPE_HEAD_DIM
V_HEAD_DIM = KV_LORA_RANK
HALF_V_DIM = V_HEAD_DIM // 2
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
LOG2E = 1.4426950408889634
SM_SCALE_LOG2E = SM_SCALE * LOG2E
FP8_DTYPE = aiter_dtypes.fp8
PAGE_SIZE_TABLE = {
(4, 1024): 1,
(32, 1024): 1,
(64, 1024): 2,
(256, 1024): 2,
(4, 8192): 8,
(32, 8192): 8,
(64, 8192): 8,
(256, 8192): 8,
}
NUM_KV_SPLITS_TABLE = {
(4, 1024): 16,
(4, 8192): 32,
(32, 1024): 8,
(32, 8192): 8,
(64, 1024): 4,
(64, 8192): 8,
(256, 1024): 1,
(256, 8192): 16,
}
MXFP4_SPLITS_TABLE = {
(32, 8192): 8, # 32 * 8 = 256 stage1 programs
(64, 8192): 4, # 64 * 4 = 256 stage1 programs
(256, 8192): 2, # 256 * 2 = 512 stage1 programs; lower stage2 traffic than 4-way split
}
SHAPE_STRATEGY = {
(4, 1024): "a8w8", (4, 8192): "a8w8", # small batch: fp8 Q saves bandwidth
(32, 1024): "a16w8", (32, 8192): "mxfp4", # large batch short kv: skip Q quant
(64, 1024): "a16w8", (64, 8192): "mxfp4",
(256, 1024): "a16w8", (256, 8192): "mxfp4",
}
DEFAULT_STRATEGY = "a16w8"
def _get_page_size(batch_size: int, kv_seq_len: int) -> int:
return PAGE_SIZE_TABLE.get((batch_size, kv_seq_len), 2)
def _get_num_kv_splits(batch_size: int, kv_seq_len: int) -> int:
return NUM_KV_SPLITS_TABLE.get((batch_size, kv_seq_len), 16)
def _get_mxfp4_splits(batch_size: int, kv_seq_len: int) -> int:
return MXFP4_SPLITS_TABLE.get((batch_size, kv_seq_len), 1)
_E2M1_LUT = None
def _get_e2m1_lut() -> torch.Tensor:
global _E2M1_LUT
if _E2M1_LUT is None:
_E2M1_LUT = torch.tensor(
[0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0,
0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0],
dtype=torch.float16,
device="cuda",
)
return _E2M1_LUT
# =====================================================================================
# MXFP4 qh16 kernel: one program handles all 16 query heads for one (request, split)
# =====================================================================================
@triton.jit
def _mla_mxfp4_qh16_stage1_reuse(
Q_packed, # [batch, H, 288] uint8
Q_scale, # [batch*H, 18] uint8 (viewed as byte rows; actual row stride >= 18)
KV_packed, # [total_kv, 288] uint8
KV_scale, # [total_kv, 18] uint8
E2M1_LUT, # [16] fp16
Mid_O, # [batch*H, S, 512] fp32, layout = [even256 | odd256]
Mid_LSE, # [batch*H, S] fp32, log2-domain lse
kv_indptr, # [batch+1] int32
sm_scale_log2e,
stride_qp_b: tl.constexpr,
stride_qp_h: tl.constexpr,
stride_qs,
stride_kv,
stride_ks,
stride_mid_bh,
stride_mid_s,
stride_lse_bh,
stride_lse_s,
NUM_HEADS: tl.constexpr,
BLOCK_N: tl.constexpr,
NUM_SPLITS: tl.constexpr,
HALF_C: tl.constexpr,
BLOCK_V: tl.constexpr,
):
pid = tl.program_id(0)
batch_id = pid // NUM_SPLITS
split_id = pid % NUM_SPLITS
kv_begin = tl.load(kv_indptr + batch_id)
seq_len = tl.load(kv_indptr + batch_id + 1) - kv_begin
kv_per_split = tl.cdiv(seq_len, NUM_SPLITS)
split_start = kv_per_split * split_id
split_end = tl.minimum(split_start + kv_per_split, seq_len)
offs_h = tl.arange(0, NUM_HEADS)
offs_v = tl.arange(0, BLOCK_V)
if split_start >= split_end:
zptrs = Mid_O + (batch_id * NUM_HEADS + offs_h[:, None]) * stride_mid_bh + split_id * stride_mid_s + offs_v[None, :]
tl.store(zptrs, tl.zeros([NUM_HEADS, BLOCK_V], dtype=tl.float32))
tl.store(zptrs + BLOCK_V, tl.zeros([NUM_HEADS, BLOCK_V], dtype=tl.float32))
tl.store(zptrs + HALF_C, tl.zeros([NUM_HEADS, BLOCK_V], dtype=tl.float32))
tl.store(zptrs + HALF_C + BLOCK_V, tl.zeros([NUM_HEADS, BLOCK_V], dtype=tl.float32))
lptrs = Mid_LSE + (batch_id * NUM_HEADS + offs_h) * stride_lse_bh + split_id * stride_lse_s
tl.store(lptrs, tl.full([NUM_HEADS], float("-inf"), dtype=tl.float32))
return
# Q layout: first 512 dims => 2 x 128 packed-byte halves; last 64 dims => 32 packed bytes.
offs_c0 = tl.arange(0, BLOCK_V)
offs_c1 = BLOCK_V + offs_c0
offs_r = tl.arange(0, 32)
q_row = batch_id * NUM_HEADS
q_nope0 = tl.load(
Q_packed + batch_id * stride_qp_b + offs_h[:, None] * stride_qp_h + offs_c0[None, :]
)
q_sc0 = tl.load(
Q_scale + (q_row + offs_h[:, None]) * stride_qs + tl.arange(0, 8)[None, :]
)
q_nope1 = tl.load(
Q_packed + batch_id * stride_qp_b + offs_h[:, None] * stride_qp_h + offs_c1[None, :]
)
q_sc1 = tl.load(
Q_scale + (q_row + offs_h[:, None]) * stride_qs + 8 + tl.arange(0, 8)[None, :]
)
q_rope = tl.load(
Q_packed + batch_id * stride_qp_b + offs_h[:, None] * stride_qp_h + HALF_C + offs_r[None, :]
)
q_sc_rope = tl.load(
Q_scale + (q_row + offs_h[:, None]) * stride_qs + 16 + tl.arange(0, 2)[None, :]
)
acc_e0 = tl.zeros([NUM_HEADS, BLOCK_V], dtype=tl.float32)
acc_e1 = tl.zeros([NUM_HEADS, BLOCK_V], dtype=tl.float32)
acc_o0 = tl.zeros([NUM_HEADS, BLOCK_V], dtype=tl.float32)
acc_o1 = tl.zeros([NUM_HEADS, BLOCK_V], dtype=tl.float32)
m_i = tl.full([NUM_HEADS], float("-inf"), dtype=tl.float32)
l_i = tl.zeros([NUM_HEADS], dtype=tl.float32)
offs_n = tl.arange(0, BLOCK_N)
sc_v_idx = offs_v // 16 # 128 bytes -> 8 scale groups
for start_n in range(split_start, split_end, BLOCK_N):
cur_n = start_n + offs_n
mask_n = cur_n < split_end
kv_locs = kv_begin + cur_n # contiguous within request; no separate kv_indices indirection
# Load first 512 dims once as two 128-byte halves. Each half is reused for both QK(nope) and PV.
k_nope0 = tl.load(
KV_packed + kv_locs[:, None] * stride_kv + offs_c0[None, :],
mask=mask_n[:, None],
other=0,
)
k_sc0 = tl.load(
KV_scale + kv_locs[:, None] * stride_ks + tl.arange(0, 8)[None, :],
mask=mask_n[:, None],
other=0,
)
k_nope1 = tl.load(
KV_packed + kv_locs[:, None] * stride_kv + offs_c1[None, :],
mask=mask_n[:, None],
other=0,
)
k_sc1 = tl.load(
KV_scale + kv_locs[:, None] * stride_ks + 8 + tl.arange(0, 8)[None, :],
mask=mask_n[:, None],
other=0,
)
qk = tl.zeros([NUM_HEADS, BLOCK_N], dtype=tl.float32)
qk = tl.dot_scaled(
q_nope0, q_sc0, "e2m1",
tl.trans(k_nope0), k_sc0, "e2m1",
acc=qk,
fast_math=True,
)
qk = tl.dot_scaled(
q_nope1, q_sc1, "e2m1",
tl.trans(k_nope1), k_sc1, "e2m1",
acc=qk,
fast_math=True,
)
k_rope = tl.load(
KV_packed + kv_locs[:, None] * stride_kv + HALF_C + offs_r[None, :],
mask=mask_n[:, None],
other=0,
)
k_sc_rope = tl.load(
KV_scale + kv_locs[:, None] * stride_ks + 16 + tl.arange(0, 2)[None, :],
mask=mask_n[:, None],
other=0,
)
qk = tl.dot_scaled(
q_rope, q_sc_rope, "e2m1",
tl.trans(k_rope), k_sc_rope, "e2m1",
acc=qk,
fast_math=True,
)
qk = qk * sm_scale_log2e
qk = tl.where(mask_n[None, :], qk, float("-inf"))
qk_max = tl.max(qk, axis=1)
m_new = tl.maximum(m_i, qk_max)
alpha = tl.exp2(m_i - m_new)
p = tl.exp2(qk - tl.reshape(m_new, [NUM_HEADS, 1]))
l_i = l_i * alpha + tl.sum(p, axis=1)
alpha_2d = tl.reshape(alpha, [NUM_HEADS, 1])
acc_e0 = acc_e0 * alpha_2d
acc_e1 = acc_e1 * alpha_2d
acc_o0 = acc_o0 * alpha_2d
acc_o1 = acc_o1 * alpha_2d
m_i = m_new
p_f16 = p.to(tl.float16)
# Reuse the already-loaded first 512 dims for PV.
v0_i = k_nope0.to(tl.int32)
sc0 = k_sc0[:, sc_v_idx]
sc0_f = tl.exp2(sc0.to(tl.float32) - 127.0)
lo0 = v0_i & 0xF
hi0 = (v0_i >> 4) & 0xF
vlo0 = tl.load(E2M1_LUT + lo0)
vhi0 = tl.load(E2M1_LUT + hi0)
vlo0_s = (vlo0.to(tl.float32) * sc0_f).to(tl.float16)
vhi0_s = (vhi0.to(tl.float32) * sc0_f).to(tl.float16)
acc_e0 += tl.dot(p_f16, vlo0_s, out_dtype=tl.float32)
acc_o0 += tl.dot(p_f16, vhi0_s, out_dtype=tl.float32)
v1_i = k_nope1.to(tl.int32)
sc1 = k_sc1[:, sc_v_idx]
sc1_f = tl.exp2(sc1.to(tl.float32) - 127.0)
lo1 = v1_i & 0xF
hi1 = (v1_i >> 4) & 0xF
vlo1 = tl.load(E2M1_LUT + lo1)
vhi1 = tl.load(E2M1_LUT + hi1)
vlo1_s = (vlo1.to(tl.float32) * sc1_f).to(tl.float16)
vhi1_s = (vhi1.to(tl.float32) * sc1_f).to(tl.float16)
acc_e1 += tl.dot(p_f16, vlo1_s, out_dtype=tl.float32)
acc_o1 += tl.dot(p_f16, vhi1_s, out_dtype=tl.float32)
safe_l = tl.where(l_i > 0.0, l_i, 1.0)
inv_l = 1.0 / safe_l
inv_l_2d = tl.reshape(inv_l, [NUM_HEADS, 1])
acc_e0 = acc_e0 * inv_l_2d
acc_e1 = acc_e1 * inv_l_2d
acc_o0 = acc_o0 * inv_l_2d
acc_o1 = acc_o1 * inv_l_2d
lse = m_i + tl.log(safe_l) * LOG2E
base_ptrs = Mid_O + (batch_id * NUM_HEADS + offs_h[:, None]) * stride_mid_bh + split_id * stride_mid_s
tl.store(base_ptrs + offs_v[None, :], acc_e0)
tl.store(base_ptrs + BLOCK_V + offs_v[None, :], acc_e1)
tl.store(base_ptrs + HALF_C + offs_v[None, :], acc_o0)
tl.store(base_ptrs + HALF_C + BLOCK_V + offs_v[None, :], acc_o1)
lptrs = Mid_LSE + (batch_id * NUM_HEADS + offs_h) * stride_lse_bh + split_id * stride_lse_s
tl.store(lptrs, lse)
@triton.jit
def _mla_mxfp4_stage2_log2(
Mid_O,
Mid_LSE,
O,
stride_mid_bh,
stride_mid_s,
stride_lse_bh,
stride_lse_s,
stride_o_bh,
HALF_C: tl.constexpr,
NUM_SPLITS: tl.constexpr,
):
pid = tl.program_id(0)
offs = tl.arange(0, HALF_C)
e_max = tl.full([], float("-inf"), dtype=tl.float32)
e_sum = tl.full([], 0.0, dtype=tl.float32)
r_even = tl.zeros([HALF_C], dtype=tl.float32)
r_odd = tl.zeros([HALF_C], dtype=tl.float32)
mid_base = Mid_O + pid * stride_mid_bh
lse_base = Mid_LSE + pid * stride_lse_bh
for s in tl.static_range(0, NUM_SPLITS):
sb = mid_base + s * stride_mid_s
sv_e = tl.load(sb + offs)
sv_o = tl.load(sb + HALF_C + offs)
lse = tl.load(lse_base + s * stride_lse_s)
n_max = tl.maximum(lse, e_max)
old_sc = tl.exp2(e_max - n_max)
el = tl.exp2(lse - n_max)
r_even = r_even * old_sc + el * sv_e
r_odd = r_odd * old_sc + el * sv_o
e_sum = e_sum * old_sc + el
e_max = n_max
inv_s = 1.0 / tl.where(e_sum > 0.0, e_sum, 1.0)
r_even = r_even * inv_s
r_odd = r_odd * inv_s
o_base = O + pid * stride_o_bh
tl.store(o_base + offs * 2, r_even.to(tl.bfloat16))
tl.store(o_base + offs * 2 + 1, r_odd.to(tl.bfloat16))
# =====================================================================================
# AITER direct stage1/reduce fallback
# =====================================================================================
def _mla_stage1_direct(
q_ready: torch.Tensor,
kv_buffer_4d: torch.Tensor,
cached: dict,
output: torch.Tensor,
q_scale,
kv_scale: torch.Tensor,
):
aiter_module.mla_decode_stage1_asm_fwd(
q_ready,
kv_buffer_4d,
cached["qo_indptr"],
cached["kv_indptr"],
cached["kv_indices"],
cached["kv_last_page_len"],
None, # num_kv_splits_indptr = None -> persistent mode
cached["work_meta_data"],
cached["work_indptr"],
cached["work_info_set"],
1,
cached["page_size"],
NUM_KV_HEADS,
SM_SCALE,
cached["logits"],
cached["attn_lse"],
output,
q_scale,
kv_scale,
)
def _mla_reduce_direct(cached: dict, output: torch.Tensor):
aiter_module.mla_reduce_v1(
cached["logits"],
cached["attn_lse"],
cached["reduce_indptr"],
cached["reduce_final_map"],
cached["reduce_partial_map"],
1,
output,
)
# =====================================================================================
# Cache builders
# =====================================================================================
_cache = {}
_decided_strategy = {}
def _build_mxfp4_cache(batch_size: int, kv_seq_len: int) -> dict:
key = ("mxfp4", batch_size, kv_seq_len)
if key in _cache:
return _cache[key]
num_splits = _get_mxfp4_splits(batch_size, kv_seq_len)
entry = {
"num_splits": num_splits,
"mid_o": torch.empty(
(batch_size * NUM_HEADS, num_splits, V_HEAD_DIM),
dtype=torch.float32,
device="cuda",
),
"mid_lse": torch.empty(
(batch_size * NUM_HEADS, num_splits),
dtype=torch.float32,
device="cuda",
),
"output": torch.empty(
(batch_size, NUM_HEADS, V_HEAD_DIM),
dtype=torch.bfloat16,
device="cuda",
),
}
_cache[key] = entry
return entry
def _build_a8w8_cache(batch_size: int, kv_seq_len: int) -> dict:
key = ("a8w8", batch_size, kv_seq_len)
if key in _cache:
return _cache[key]
total_q = batch_size
num_kv_splits = _get_num_kv_splits(batch_size, kv_seq_len)
page_size = _get_page_size(batch_size, kv_seq_len)
num_pages_per_batch = kv_seq_len // page_size
total_pages = batch_size * num_pages_per_batch
qo_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda")
kv_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda") * num_pages_per_batch
kv_last_page_len = torch.full((batch_size,), page_size, dtype=torch.int32, device="cuda")
kv_indices = torch.arange(total_pages, dtype=torch.int32, device="cuda")
info = get_mla_metadata_info_v1(
batch_size,
1,
NUM_HEADS,
FP8_DTYPE,
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="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,
NUM_HEADS // NUM_KV_HEADS,
NUM_KV_HEADS,
True,
work_metadata,
work_info_set,
work_indptr,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
page_size=page_size,
kv_granularity=max(page_size, 16),
max_seqlen_qo=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=FP8_DTYPE,
)
num_partials = reduce_partial_map.size(0)
entry = {
"page_size": page_size,
"num_kv_splits": num_kv_splits,
"qo_indptr": qo_indptr,
"kv_indptr": kv_indptr,
"kv_indices": kv_indices,
"kv_last_page_len": kv_last_page_len,
"output": torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
"logits": torch.empty((num_partials, 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device="cuda"),
"attn_lse": torch.empty((num_partials, 1, NUM_HEADS, 1), dtype=torch.float32, device="cuda"),
"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,
"q_fp8_buf": torch.empty((total_q, NUM_HEADS, QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda"),
"q_scale_buf": torch.empty(1, dtype=torch.float32, device="cuda"),
}
_cache[key] = entry
return entry
def _build_a16w8_cache(batch_size: int, kv_seq_len: int) -> dict:
key = ("a16w8", batch_size, kv_seq_len)
if key in _cache:
return _cache[key]
total_q = batch_size
num_kv_splits = _get_num_kv_splits(batch_size, kv_seq_len)
page_size = _get_page_size(batch_size, kv_seq_len)
num_pages_per_batch = kv_seq_len // page_size
total_pages = batch_size * num_pages_per_batch
qo_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda")
kv_indptr = torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda") * num_pages_per_batch
kv_last_page_len = torch.full((batch_size,), page_size, dtype=torch.int32, device="cuda")
kv_indices = torch.arange(total_pages, dtype=torch.int32, device="cuda")
q_dtype = torch.bfloat16
kv_dtype = FP8_DTYPE
info = get_mla_metadata_info_v1(
batch_size, 1, NUM_HEADS, 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,
NUM_HEADS // NUM_KV_HEADS, NUM_KV_HEADS, True,
work_metadata, work_info_set, work_indptr,
reduce_indptr, reduce_final_map, reduce_partial_map,
page_size=page_size, kv_granularity=max(page_size, 16),
max_seqlen_qo=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,
)
num_partials = reduce_partial_map.size(0)
entry = {
"page_size": page_size,
"num_kv_splits": num_kv_splits,
"qo_indptr": qo_indptr,
"kv_indptr": kv_indptr,
"kv_indices": kv_indices,
"kv_last_page_len": kv_last_page_len,
"output": torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
"logits": torch.empty((num_partials, 1, NUM_HEADS, V_HEAD_DIM), dtype=torch.float32, device="cuda"),
"attn_lse": torch.empty((num_partials, 1, NUM_HEADS, 1), dtype=torch.float32, device="cuda"),
"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,
}
_cache[key] = entry
return entry
# =====================================================================================
# Runtime helpers
# =====================================================================================
def _quantize_fp8_naive(tensor: torch.Tensor):
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)
def _run_a8w8_direct(q_bf16: torch.Tensor, kv_buffer_4d: torch.Tensor, kv_scale: torch.Tensor, cached: dict) -> torch.Tensor:
if HAS_FUSED_QUANT:
try:
dynamic_per_tensor_quant(cached["q_fp8_buf"], q_bf16.view_as(cached["q_fp8_buf"]), cached["q_scale_buf"])
q_fp8 = cached["q_fp8_buf"]
q_scale = cached["q_scale_buf"]
except Exception:
q_fp8, q_scale = _quantize_fp8_naive(q_bf16)
else:
q_fp8, q_scale = _quantize_fp8_naive(q_bf16)
output = cached["output"]
_mla_stage1_direct(q_fp8.view(-1, NUM_HEADS, QK_HEAD_DIM), kv_buffer_4d, cached, output, q_scale, kv_scale)
_mla_reduce_direct(cached, output)
return output
def _run_a16w8_direct(q_bf16: torch.Tensor, kv_buffer_4d: torch.Tensor, kv_scale: torch.Tensor, cached: dict) -> torch.Tensor:
output = cached["output"]
_mla_stage1_direct(q_bf16.view(-1, NUM_HEADS, QK_HEAD_DIM), kv_buffer_4d, cached, output, None, kv_scale)
_mla_reduce_direct(cached, output)
return output
def _run_mxfp4_qh16(q_bf16: torch.Tensor, kv_data_mxfp4, kv_indptr: torch.Tensor, cached: dict) -> torch.Tensor:
kv_packed, kv_scale = kv_data_mxfp4
num_splits = cached["num_splits"]
q_2d = q_bf16.view(-1, QK_HEAD_DIM)
q_packed_2d, q_scale_2d = dynamic_mxfp4_quant(q_2d)
q_packed_u8 = q_packed_2d.view(torch.uint8)
q_scale_u8 = q_scale_2d.view(torch.uint8)
q_packed_3d = q_packed_u8.reshape(q_bf16.shape[0], NUM_HEADS, QK_HEAD_DIM // 2)
total_kv = kv_packed.shape[0]
kv_packed_2d = kv_packed.reshape(total_kv, -1).view(torch.uint8)
kv_scale_2d = kv_scale.reshape(total_kv, -1).view(torch.uint8)
mid_o = cached["mid_o"]
mid_lse = cached["mid_lse"]
output = cached["output"]
output_flat = output.view(-1, V_HEAD_DIM)
lut = _get_e2m1_lut()
BLOCK_N = 32
BLOCK_V = 128
grid1 = (q_bf16.shape[0] * num_splits,)
_mla_mxfp4_qh16_stage1_reuse[grid1](
q_packed_3d,
q_scale_u8,
kv_packed_2d,
kv_scale_2d,
lut,
mid_o,
mid_lse,
kv_indptr,
SM_SCALE_LOG2E,
q_packed_3d.stride(0),
q_packed_3d.stride(1),
q_scale_u8.stride(0),
kv_packed_2d.stride(0),
kv_scale_2d.stride(0),
mid_o.stride(0),
mid_o.stride(1),
mid_lse.stride(0),
mid_lse.stride(1),
NUM_HEADS=NUM_HEADS,
BLOCK_N=BLOCK_N,
NUM_SPLITS=num_splits,
HALF_C=HALF_V_DIM,
BLOCK_V=BLOCK_V,
num_warps=4,
num_stages=1,
)
grid2 = (q_bf16.shape[0] * NUM_HEADS,)
_mla_mxfp4_stage2_log2[grid2](
mid_o,
mid_lse,
output_flat,
mid_o.stride(0),
mid_o.stride(1),
mid_lse.stride(0),
mid_lse.stride(1),
output_flat.stride(0),
HALF_C=HALF_V_DIM,
NUM_SPLITS=num_splits,
num_warps=4,
num_stages=1,
)
return output
# =====================================================================================
# Reference and custom entrypoints
# =====================================================================================
def ref_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
kv_seq_len = config.get("kv_seq_len", 1024)
q_input, q_scale = _quantize_fp8_naive(q)
kv_input, kv_scale = kv_data["fp8"]
page_size = _get_page_size(batch_size, kv_seq_len)
total_pages = kv_input.shape[0] // page_size
kv_buffer_4d = kv_input.view(total_pages, page_size, NUM_KV_HEADS, kv_input.shape[-1])
num_pages_per_batch = kv_seq_len // page_size
kv_indptr_paged = torch.arange(0, batch_size + 1, dtype=torch.int32, device="cuda") * num_pages_per_batch
kv_indices_ref = torch.arange(total_pages, dtype=torch.int32, device="cuda")
kv_last_page_len = torch.full((batch_size,), page_size, dtype=torch.int32, device="cuda")
num_kv_splits = _get_num_kv_splits(batch_size, kv_seq_len)
info = get_mla_metadata_info_v1(
batch_size,
1,
NUM_HEADS,
q_input.dtype,
kv_input.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]
(wm, wi, wis, ri, rfm, rpm) = work
get_mla_metadata_v1(
qo_indptr,
kv_indptr_paged,
kv_last_page_len,
NUM_HEADS // NUM_KV_HEADS,
NUM_KV_HEADS,
True,
wm,
wis,
wi,
ri,
rfm,
rpm,
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=num_kv_splits,
intra_batch_mode=True,
dtype_q=q_input.dtype,
dtype_kv=kv_input.dtype,
)
output = torch.empty((q.shape[0], NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
mla_decode_fwd(
q_input.view(-1, NUM_HEADS, QK_HEAD_DIM),
kv_buffer_4d,
output,
qo_indptr,
kv_indptr_paged,
kv_indices_ref,
kv_last_page_len,
1,
page_size=page_size,
nhead_kv=NUM_KV_HEADS,
sm_scale=SM_SCALE,
logit_cap=0.0,
num_kv_splits=num_kv_splits,
q_scale=q_scale,
kv_scale=kv_scale,
intra_batch_mode=True,
work_meta_data=wm,
work_indptr=wi,
work_info_set=wis,
reduce_indptr=ri,
reduce_final_map=rfm,
reduce_partial_map=rpm,
)
return output
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
batch_size = config["batch_size"]
kv_seq_len = config["kv_seq_len"]
shape_key = (batch_size, kv_seq_len)
strategy = _decided_strategy.get(shape_key)
if strategy is None:
strategy = SHAPE_STRATEGY.get(shape_key, DEFAULT_STRATEGY)
if strategy == "mxfp4" and not HAS_MXFP4:
strategy = "a16w8"
if strategy == "mxfp4":
try:
cached = _build_mxfp4_cache(batch_size, kv_seq_len)
result = _run_mxfp4_qh16(q, kv_data["mxfp4"], kv_indptr, cached)
_decided_strategy[shape_key] = "mxfp4"
return result
except Exception:
strategy = "a16w8"
kv_input, kv_scale = kv_data["fp8"]
page_size = _get_page_size(batch_size, kv_seq_len)
kv_buffer_4d = kv_input.view(
kv_input.shape[0] // page_size,
page_size,
NUM_KV_HEADS,
kv_input.shape[-1],
)
if strategy == "a16w8":
cached = _build_a16w8_cache(batch_size, kv_seq_len)
result = _run_a16w8_direct(q, kv_buffer_4d, kv_scale, cached)
_decided_strategy[shape_key] = "a16w8"
return result
cached = _build_a8w8_cache(batch_size, kv_seq_len)
result = _run_a8w8_direct(q, kv_buffer_4d, kv_scale, cached)
_decided_strategy[shape_key] = "a8w8"
return result
scrolls · 838 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