submission 587387
garrick99 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 396 lines, June 9 Researcher Reciprocity License v1.0.
submission_v4_dotscaled.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-587387?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:9c75d95dc4fce8f370c89eaf39643ede33492d02b3dc9a1c9fb92d4d85a82fbd
license declaredunknown
license concludedunknown
authorsgarrick99
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
K scores use native fp4 MFMA via dot_scaled (eliminates 9-tile dequant overhead).mma
acc_even += tl.dot(p_bf, ve).to(tl.float32)online-softmax
m_new = tl.maximum(m_i, blk_max)Kernel source
submission_v4_dotscaled.py396 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""
MLA Decode v4 — tl.dot_scaled for K scores + manual dequant for V.
K scores use native fp4 MFMA via dot_scaled (eliminates 9-tile dequant overhead).
V accumulation uses even/odd manual dequant (packed axis is N, not K for V).
Falls back to aiter fp8 if dot_scaled compilation fails.
"""
import torch
import triton
import triton.language as tl
import math
from task import input_t, output_t
# ============================================================================
# FP4 dequant helpers (for V accumulation only)
# ============================================================================
@triton.jit
def _fp4_to_f32(nibble):
sign = (nibble >> 3) & 1
abs3 = nibble & 7
exp2 = abs3 >> 1
man1 = abs3 & 1
val = tl.where(
exp2 > 0,
tl.math.exp2((exp2 - 1).to(tl.float32)) * (1.0 + man1.to(tl.float32) * 0.5),
man1.to(tl.float32) * 0.5,
)
return tl.where(sign > 0, -val, val)
@triton.jit
def _e8m0_to_f32(byte_val):
return tl.math.exp2(byte_val.to(tl.float32) - 127.0)
# ============================================================================
# Stage 1: dot_scaled K scores + manual V dequant
# ============================================================================
@triton.jit
def _mla_v4_stage1(
QT_ptr, # (B, D_QK, H) bf16 — Q transposed
QE_ptr, QO_ptr, # (B, H, half_dk) bf16 — Q even/odd (unused if dot_scaled works for K)
KV_ptr, # (total_kv, packed_dim) uint8
KS_ptr, # (scale_rows, scale_cols) uint8
POE_ptr, POO_ptr, PLSE_ptr,
KVIDX_ptr,
sm_scale,
# QT strides
sqt_b, sqt_k, sqt_h,
# QE/QO strides (for V path's Q — not needed for K with dot_scaled)
sqe_b, sqe_h, sqe_d,
sqo_b, sqo_h, sqo_d,
# KV packed strides
skv_t, skv_d,
# KV scale strides
sks_t, sks_d,
# POE strides
spe_b, spe_s, spe_h, spe_d,
# POO strides
spo_b, spo_s, spo_h, spo_d,
# PLSE strides
sl_b, sl_s, sl_h,
# constexpr
HALF_DV: tl.constexpr, # 256
N_HEADS: tl.constexpr, # 16
NUM_SPLITS: tl.constexpr,
BLOCK_KV: tl.constexpr, # 64
K_TILE_PACKED: tl.constexpr, # 32
K_TILE: tl.constexpr, # 64
N_K_TILES: tl.constexpr, # 9
):
bid = tl.program_id(0)
sid = tl.program_id(1)
kv_start = tl.load(KVIDX_ptr + bid)
kv_end = tl.load(KVIDX_ptr + bid + 1)
kv_len = kv_end - kv_start
tps = tl.cdiv(kv_len, NUM_SPLITS)
s_start = sid * tps
s_end = tl.minimum(s_start + tps, kv_len)
hr = tl.arange(0, N_HEADS) # 16
tp = tl.arange(0, K_TILE_PACKED) # 32
tk = tl.arange(0, K_TILE) # 64
bk = tl.arange(0, BLOCK_KV) # 64
dv = tl.arange(0, HALF_DV) # 256
sc2 = tl.arange(0, 2) # 2 scale blocks per K tile
if s_start >= kv_len:
tl.store(PLSE_ptr + bid * sl_b + sid * sl_s + hr * sl_h,
tl.full([N_HEADS], float('-inf'), tl.float32))
return
# Initialize
m_i = tl.full([N_HEADS], float('-inf'), tl.float32)
l_i = tl.zeros([N_HEADS], tl.float32)
acc_even = tl.zeros([N_HEADS, HALF_DV], tl.float32)
acc_odd = tl.zeros([N_HEADS, HALF_DV], tl.float32)
for blk_off in range(s_start, s_end, BLOCK_KV):
blk_len = tl.minimum(BLOCK_KV, s_end - blk_off)
vmask = bk < blk_len
tokens = kv_start + blk_off + bk
# ======== K scores via tl.dot_scaled (native fp4 MFMA) ========
scores_t = tl.zeros([BLOCK_KV, N_HEADS], tl.float32)
for kt in tl.static_range(0, N_K_TILES):
tile_off_p = kt * K_TILE_PACKED
tile_off_l = kt * K_TILE
# K tile: (BK, 32_packed) fp4
kp_off = tokens[:, None] * skv_t + (tile_off_p + tp)[None, :] * skv_d
k_tile = tl.load(KV_ptr + kp_off, mask=vmask[:, None], other=0)
# K scale: (BK, 2) E8M0
ks_off = tokens[:, None] * sks_t + (kt * 2 + sc2)[None, :] * sks_d
k_scale = tl.load(KS_ptr + ks_off, mask=vmask[:, None], other=0)
# Q^T tile: (K_TILE=64, H=16) bf16
qt_off = bid * sqt_b + (tile_off_l + tk)[:, None] * sqt_k + hr[None, :] * sqt_h
qt_tile = tl.load(QT_ptr + qt_off)
# dot_scaled: K_fp4(BK,32packed) @ Q^T_fp16(64,16) -> (BK,16)
scores_t = tl.dot_scaled(
lhs=k_tile,
rhs=qt_tile.to(tl.float16),
lhs_scale=k_scale,
rhs_scale=None,
lhs_format='e2m1',
rhs_format='fp16',
acc=scores_t,
)
# Transpose: (BK, H) -> (H, BK) and apply scale
scores = tl.trans(scores_t) * sm_scale
scores = tl.where(vmask[None, :], scores, float('-inf'))
# ======== Online softmax ========
blk_max = tl.max(scores, axis=1)
m_new = tl.maximum(m_i, blk_max)
alpha = tl.exp(m_i - m_new)
p = tl.exp(scores - m_new[:, None])
l_i = l_i * alpha + tl.sum(p, axis=1)
acc_even = acc_even * alpha[:, None]
acc_odd = acc_odd * alpha[:, None]
m_i = m_new
# ======== V accumulation: manual fp4 dequant (even/odd) ========
vp_off = tokens[:, None] * skv_t + dv[None, :] * skv_d
packed_v = tl.load(KV_ptr + vp_off, mask=vmask[:, None], other=0).to(tl.int32)
v_lo = _fp4_to_f32(packed_v & 0x0F)
v_hi = _fp4_to_f32((packed_v >> 4) & 0x0F)
vs_off = tokens[:, None] * sks_t + (dv[None, :] // 16) * sks_d
v_scale = _e8m0_to_f32(
tl.load(KS_ptr + vs_off, mask=vmask[:, None], other=0).to(tl.int32))
ve = (v_lo * v_scale).to(tl.bfloat16)
vo = (v_hi * v_scale).to(tl.bfloat16)
p_bf = p.to(tl.bfloat16)
acc_even += tl.dot(p_bf, ve).to(tl.float32)
acc_odd += tl.dot(p_bf, vo).to(tl.float32)
# Store
oe = acc_even / l_i[:, None]
oo = acc_odd / l_i[:, None]
lse = m_i + tl.log(l_i)
pe_base = bid * spe_b + sid * spe_s
tl.store(POE_ptr + pe_base + hr[:, None] * spe_h + dv[None, :] * spe_d, oe)
po_base = bid * spo_b + sid * spo_s
tl.store(POO_ptr + po_base + hr[:, None] * spo_h + dv[None, :] * spo_d, oo)
tl.store(PLSE_ptr + bid * sl_b + sid * sl_s + hr * sl_h, lse)
# ============================================================================
# Stage 2: Reduce + interleave (same as v3)
# ============================================================================
@triton.jit
def _mla_reduce(
POE_ptr, POO_ptr, PLSE_ptr, O_ptr,
spe_b, spe_s, spe_h, spe_d,
spo_b, spo_s, spo_h, spo_d,
sl_b, sl_s, sl_h,
so_b, so_h, so_d,
NUM_SPLITS: tl.constexpr,
HALF_V: tl.constexpr,
):
bid = tl.program_id(0)
hid = tl.program_id(1)
dv = tl.arange(0, HALF_V)
max_lse = tl.full([], float('-inf'), tl.float32)
for s in tl.static_range(0, NUM_SPLITS):
lse = tl.load(PLSE_ptr + bid * sl_b + s * sl_s + hid * sl_h)
max_lse = tl.maximum(max_lse, lse)
acc_e = tl.zeros([HALF_V], tl.float32)
acc_o = tl.zeros([HALF_V], tl.float32)
sum_w = tl.full([], 0.0, tl.float32)
for s in tl.static_range(0, NUM_SPLITS):
lse = tl.load(PLSE_ptr + bid * sl_b + s * sl_s + hid * sl_h)
w = tl.exp(lse - max_lse)
sum_w += w
acc_e += w * tl.load(POE_ptr + bid * spe_b + s * spe_s + hid * spe_h + dv * spe_d)
acc_o += w * tl.load(POO_ptr + bid * spo_b + s * spo_s + hid * spo_h + dv * spo_d)
acc_e = (acc_e / sum_w).to(tl.bfloat16)
acc_o = (acc_o / sum_w).to(tl.bfloat16)
o_base = bid * so_b + hid * so_h
tl.store(O_ptr + o_base + (dv * 2) * so_d, acc_e)
tl.store(O_ptr + o_base + (dv * 2 + 1) * so_d, acc_o)
# ============================================================================
# Aiter FP8 fallback
# ============================================================================
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
FP8_DTYPE = aiter_dtypes.fp8
_SM_SCALE = 1.0 / (576 ** 0.5)
_aiter_meta_cache = {}
_aiter_kvidx_cache = {}
def _quantize_fp8(tensor):
finfo = torch.finfo(FP8_DTYPE)
amax = tensor.abs().amax().clamp(min=1e-12)
scale = amax / finfo.max
fp8_t = (tensor / scale).clamp(min=finfo.min, max=finfo.max).to(FP8_DTYPE)
return fp8_t, scale.to(torch.float32).reshape(1)
def _aiter_fp8_path(q, kv_data, qo_indptr, kv_indptr, config):
B = config["batch_size"]; nq = config["num_heads"]; nkv = config["num_kv_heads"]
dq = config["qk_head_dim"]; dv = config["v_head_dim"]; qsl = config["q_seq_len"]
nks = 32
q_fp8, q_scale = _quantize_fp8(q)
kv_fp8, kv_scale = kv_data["fp8"]
total_kv = int(kv_indptr[-1].item())
if total_kv not in _aiter_kvidx_cache:
_aiter_kvidx_cache[total_kv] = torch.arange(total_kv, dtype=torch.int32, device="cuda")
kv_indices = _aiter_kvidx_cache[total_kv]
kv_4d = kv_fp8.view(kv_fp8.shape[0], 1, nkv, kv_fp8.shape[-1])
kv_last = (kv_indptr[1:] - kv_indptr[:-1]).to(torch.int32)
meta_key = (B, config["kv_seq_len"])
if meta_key not in _aiter_meta_cache:
info = get_mla_metadata_info_v1(B, qsl, nq, q_fp8.dtype, kv_fp8.dtype,
is_sparse=False, fast_mode=False, num_kv_splits=nks, 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_last, nq // nkv, nkv, True,
wm, wis, wi, ri, rfm, rpm, page_size=1, kv_granularity=16,
max_seqlen_qo=qsl, uni_seqlen_qo=qsl, fast_mode=False,
max_split_per_batch=nks, intra_batch_mode=True,
dtype_q=q_fp8.dtype, dtype_kv=kv_fp8.dtype)
_aiter_meta_cache[meta_key] = (wm, wi, wis, ri, rfm, rpm)
wm, wi, wis, ri, rfm, rpm = _aiter_meta_cache[meta_key]
o = torch.empty((q.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, kv_indices,
kv_last, qsl, page_size=1, nhead_kv=nkv, sm_scale=_SM_SCALE, logit_cap=0.0,
num_kv_splits=nks, 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 o
# ============================================================================
# dot_scaled Triton path
# ============================================================================
def _dotscaled_triton_path(q, kv_data, qo_indptr, kv_indptr, config):
B = config["batch_size"]
H = config["num_heads"]
D_QK = config["qk_head_dim"]
D_V = config["v_head_dim"]
kv_len = config["kv_seq_len"]
sm_scale = config["sm_scale"]
kv_buf, kv_scl = kv_data["mxfp4"]
total_kv = kv_buf.shape[0]
kv_flat = kv_buf.reshape(total_kv, -1)
if kv_flat.dtype != torch.uint8:
kv_flat = kv_flat.view(torch.uint8)
kv_sc = kv_scl
if kv_sc.dtype != torch.uint8:
kv_sc = kv_sc.view(torch.uint8)
# Q transposed for dot_scaled K scores
q_t = q.transpose(-2, -1).contiguous() # (B, D_QK, H)
# Q even/odd for manual V dequant
q_even = q[:, :, 0::2].contiguous()
q_odd = q[:, :, 1::2].contiguous()
BLOCK_KV = 64
K_TILE_PACKED = 32
K_TILE = 64
half_dk = D_QK // 2 # 288
half_dv = D_V // 2 # 256
n_k_tiles = half_dk // K_TILE_PACKED # 9
max_sp = max(1, kv_len // BLOCK_KV)
desired = max(1, 608 // B)
ns = min(max_sp, desired, 64)
if ns > 1:
ns = 2 ** int(math.log2(ns))
po_e = torch.empty((B, ns, H, half_dv), dtype=torch.float32, device='cuda')
po_o = torch.empty((B, ns, H, half_dv), dtype=torch.float32, device='cuda')
plse = torch.empty((B, ns, H), dtype=torch.float32, device='cuda')
_mla_v4_stage1[(B, ns)](
q_t, q_even, q_odd,
kv_flat, kv_sc,
po_e, po_o, plse,
kv_indptr,
sm_scale,
q_t.stride(0), q_t.stride(1), q_t.stride(2),
q_even.stride(0), q_even.stride(1), q_even.stride(2),
q_odd.stride(0), q_odd.stride(1), q_odd.stride(2),
kv_flat.stride(0), kv_flat.stride(1),
kv_sc.stride(0), kv_sc.stride(1),
po_e.stride(0), po_e.stride(1), po_e.stride(2), po_e.stride(3),
po_o.stride(0), po_o.stride(1), po_o.stride(2), po_o.stride(3),
plse.stride(0), plse.stride(1), plse.stride(2),
HALF_DV=half_dv,
N_HEADS=H,
NUM_SPLITS=ns,
BLOCK_KV=BLOCK_KV,
K_TILE_PACKED=K_TILE_PACKED,
K_TILE=K_TILE,
N_K_TILES=n_k_tiles,
)
out = torch.empty((B, H, D_V), dtype=torch.bfloat16, device='cuda')
_mla_reduce[(B, H)](
po_e, po_o, plse, out,
po_e.stride(0), po_e.stride(1), po_e.stride(2), po_e.stride(3),
po_o.stride(0), po_o.stride(1), po_o.stride(2), po_o.stride(3),
plse.stride(0), plse.stride(1), plse.stride(2),
out.stride(0), out.stride(1), out.stride(2),
NUM_SPLITS=ns,
HALF_V=half_dv,
)
return out
# ============================================================================
# Entry point
# ============================================================================
_use_dotscaled = None
def custom_kernel(data: input_t) -> output_t:
global _use_dotscaled
q, kv_data, qo_indptr, kv_indptr, config = data
B = config["batch_size"]
kv_len = config["kv_seq_len"]
total = B * kv_len
if _use_dotscaled is None:
try:
result = _dotscaled_triton_path(q, kv_data, qo_indptr, kv_indptr, config)
_use_dotscaled = True
return result
except Exception as e:
import sys, traceback
print(f"dot_scaled FAILED: {type(e).__name__}: {e}", file=sys.stderr)
traceback.print_exc(file=sys.stderr)
_use_dotscaled = False
return _aiter_fp8_path(q, kv_data, qo_indptr, kv_indptr, config)
if _use_dotscaled:
if total <= 65536:
return _dotscaled_triton_path(q, kv_data, qo_indptr, kv_indptr, config)
else:
return _aiter_fp8_path(q, kv_data, qo_indptr, kv_indptr, config)
else:
return _aiter_fp8_path(q, kv_data, qo_indptr, kv_indptr, config)
scrolls · 396 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