submission 687550
rosehulman. · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1304 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-687550?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:6bd9ba7b702ce27669bc2c54ade2da9a08d791e78eccbc56985a352bc1a59ff0
license declaredunknown
license concludedunknown
authorsrosehulman.
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MLA Decode — Aiter A8W8 primary with Triton MXFP4 for bandwidth-bound cases.fp8
scores += tl.dot_scaled(q_tile, None, "e4m3", k_t, k_scale, "e2m1")mma
scores += tl.dot(q_chunk, tl.trans(k_fp8.to(tl.bfloat16)))num-warps = 4
num_warps=4, num_stages=2,online-softmax
m_new = tl.max(scores, axis=1)split-k
def _mla_splitk_small(stages = 2
num_warps=4, num_stages=2,Kernel source
submission.py1304 lines
"""
MLA Decode — Aiter A8W8 primary with Triton MXFP4 for bandwidth-bound cases.
Key design:
- Aiter A8W8 remains the lowest-risk fast path for general configs.
- Triton MXFP4 path targets large kvsl=8192 cases where HBM traffic dominates.
- qsl>1: merge q_pos into head dimension on Triton paths to reuse KV reads.
- Module-level warmup pre-compiles the Triton variants that dispatch may use.
"""
import torch
import sys
import traceback
import triton
import triton.language as tl
from task import input_t, output_t
try:
from aiter import dtypes as aiter_dtypes
from aiter.mla import mla_decode_fwd
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
import aiter as _aiter_mod
import aiter.mla as _aiter_mla_mod
_aiter_available = True
_stage1_fn = getattr(_aiter_mla_mod, 'mla_decode_stage1_asm_fwd',
getattr(_aiter_mod, 'mla_decode_stage1_asm_fwd', None))
_reduce_fn = getattr(_aiter_mla_mod, 'mla_reduce_v1',
getattr(_aiter_mod, 'mla_reduce_v1', None))
_direct_available = _stage1_fn is not None and _reduce_fn is not None
except Exception:
_aiter_available = False
_direct_available = False
QK_DIM = 576
V_DIM = 512
FP8 = torch.float8_e4m3fn
LOG2E = tl.constexpr(1.4426950408889634)
@triton.jit
def _mla_fused(
Q, KV_FP8,
pO, pM, pL, Out,
qo_indptr, kv_indptr, done_counter,
sm_scale, kv_fp8_scale,
stride_qt, stride_qh,
stride_kv,
stride_pO_s, stride_pO_t, stride_pO_h,
stride_pm_s, stride_pm_t,
stride_out_t, stride_out_h,
gen_target, nheads, qseqlen, n_hg,
NSPLITS: tl.constexpr,
BS: tl.constexpr, HPB: tl.constexpr, DK: tl.constexpr,
):
batch_id = tl.program_id(0)
split_id = tl.program_id(1)
hg_id = tl.program_id(2)
h_start = hg_id * HPB
h_off = h_start + tl.arange(0, HPB)
h_mask = h_off < nheads
q_s = tl.load(qo_indptr + batch_id)
kv_s = tl.load(kv_indptr + batch_id)
kv_e = tl.load(kv_indptr + batch_id + 1)
kv_len = kv_e - kv_s
chunk = tl.cdiv(kv_len, NSPLITS)
c_s = kv_s + split_id * chunk
c_e = tl.minimum(c_s + chunk, kv_e)
s_idx = tl.arange(0, BS)
dk_range = tl.arange(0, DK)
N_FULL_K: tl.constexpr = 576 // DK
REMAIN_K: tl.constexpr = 576 - N_FULL_K * DK
N_FULL_V: tl.constexpr = 512 // DK
combined_scale = sm_scale * kv_fp8_scale
q_pos = 0
while q_pos < qseqlen:
qtok = q_s + q_pos
mi = tl.full([HPB], float('-inf'), dtype=tl.float32)
li = tl.zeros([HPB], dtype=tl.float32)
a0 = tl.zeros([HPB, DK], dtype=tl.float32)
a1 = tl.zeros([HPB, DK], dtype=tl.float32)
a2 = tl.zeros([HPB, DK], dtype=tl.float32)
a3 = tl.zeros([HPB, DK], dtype=tl.float32)
for t_s in range(c_s, c_e, BS):
t_e = tl.minimum(t_s + BS, c_e)
smask = s_idx < (t_e - t_s)
kv_ids = t_s + s_idx
scores = tl.zeros([HPB, BS], dtype=tl.float32)
for dc in range(N_FULL_K):
q_ptrs = Q + qtok * stride_qt + h_off[:, None] * stride_qh + dc * DK + dk_range[None, :]
q_chunk = tl.load(q_ptrs, mask=h_mask[:, None], other=0.0)
k_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + dc * DK + dk_range[None, :]
k_fp8 = tl.load(k_ptrs, mask=smask[:, None], other=0.0)
scores += tl.dot(q_chunk, tl.trans(k_fp8.to(tl.bfloat16)))
if REMAIN_K > 0:
r_dk = tl.arange(0, 64)
q_ptrs = Q + qtok * stride_qt + h_off[:, None] * stride_qh + N_FULL_K * DK + r_dk[None, :]
q_chunk = tl.load(q_ptrs, mask=h_mask[:, None], other=0.0)
k_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + N_FULL_K * DK + r_dk[None, :]
k_fp8 = tl.load(k_ptrs, mask=smask[:, None], other=0.0)
scores += tl.dot(q_chunk, tl.trans(k_fp8.to(tl.bfloat16)))
scores = scores * combined_scale
scores = tl.where(smask[None, :], scores, float('-inf'))
m_new = tl.max(scores, axis=1)
m_max = tl.maximum(mi, m_new)
exp_old = tl.math.exp2((mi - m_max) * LOG2E)
exp_s = tl.math.exp2((scores - m_max[:, None]) * LOG2E)
l_tile = tl.sum(exp_s, axis=1)
l_comb = li * exp_old + l_tile
safe_l = tl.where(l_comb > 0.0, l_comb, 1.0)
r = (li * exp_old / safe_l)[:, None]
a0 *= r; a1 *= r; a2 *= r; a3 *= r
probs = exp_s / safe_l[:, None]
pb = probs.to(tl.bfloat16)
for vc in range(N_FULL_V):
v_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + vc * DK + dk_range[None, :]
v_fp8 = tl.load(v_ptrs, mask=smask[:, None], other=0.0)
c = tl.dot(pb, v_fp8.to(tl.bfloat16))
if vc == 0: a0 += c
elif vc == 1: a1 += c
elif vc == 2: a2 += c
else: a3 += c
mi = m_max
li = l_comb
a0 *= kv_fp8_scale; a1 *= kv_fp8_scale; a2 *= kv_fp8_scale; a3 *= kv_fp8_scale
po_base = pO + split_id * stride_pO_s + qtok * stride_pO_t
for vc in range(N_FULL_V):
if vc == 0: w = a0
elif vc == 1: w = a1
elif vc == 2: w = a2
else: w = a3
tl.store(po_base + h_off[:, None] * stride_pO_h + vc * DK + dk_range[None, :],
w, mask=h_mask[:, None])
pm_off = pM + split_id * stride_pm_s + qtok * stride_pm_t + h_off
pl_off = pL + split_id * stride_pm_s + qtok * stride_pm_t + h_off
tl.store(pm_off, mi, mask=h_mask)
tl.store(pl_off, li, mask=h_mask)
q_pos += 1
counter_idx = batch_id * n_hg + hg_id
old_val = tl.atomic_add(done_counter + counter_idx, 1)
is_last = (old_val + 1) == gen_target
if is_last:
N_FULL_V_R: tl.constexpr = 512 // DK
q_pos2 = 0
while q_pos2 < qseqlen:
qtok = q_s + q_pos2
gm = tl.full([HPB], float('-inf'), dtype=tl.float32)
gl = tl.zeros([HPB], dtype=tl.float32)
ra0 = tl.zeros([HPB, DK], dtype=tl.float32)
ra1 = tl.zeros([HPB, DK], dtype=tl.float32)
ra2 = tl.zeros([HPB, DK], dtype=tl.float32)
ra3 = tl.zeros([HPB, DK], dtype=tl.float32)
for s in range(NSPLITS):
pm_val = tl.load(pM + s * stride_pm_s + qtok * stride_pm_t + h_off, mask=h_mask, other=float('-inf'))
pl_val = tl.load(pL + s * stride_pm_s + qtok * stride_pm_t + h_off, mask=h_mask, other=0.0)
m_new = tl.maximum(gm, pm_val)
exp_old = tl.math.exp2((gm - m_new) * LOG2E)
exp_new = tl.math.exp2((pm_val - m_new) * LOG2E)
gl_new = gl * exp_old + pl_val * exp_new
safe_gl = tl.where(gl_new > 0.0, gl_new, 1.0)
r = (gl * exp_old / safe_gl)[:, None]
ra0 *= r; ra1 *= r; ra2 *= r; ra3 *= r
f = (pl_val * exp_new / safe_gl)[:, None]
po_base = pO + s * stride_pO_s + qtok * stride_pO_t + h_off[:, None] * stride_pO_h
for vc in range(N_FULL_V_R):
partial = tl.load(po_base + vc * DK + dk_range[None, :], mask=h_mask[:, None], other=0.0)
c = partial * f
if vc == 0: ra0 += c
elif vc == 1: ra1 += c
elif vc == 2: ra2 += c
else: ra3 += c
gm = m_new
gl = gl_new
out_base = Out + qtok * stride_out_t + h_off[:, None] * stride_out_h
for vc in range(N_FULL_V_R):
if vc == 0: w = ra0
elif vc == 1: w = ra1
elif vc == 2: w = ra2
else: w = ra3
tl.store(out_base + vc * DK + dk_range[None, :], w.to(tl.bfloat16), mask=h_mask[:, None])
q_pos2 += 1
@triton.jit
def _mla_nosplit(
Q, KV_FP8,
Out,
qo_indptr, kv_indptr,
sm_scale, kv_fp8_scale,
stride_qt, stride_qh,
stride_kv,
stride_out_t, stride_out_h,
nheads, qseqlen,
BS: tl.constexpr, HPB: tl.constexpr, DK: tl.constexpr,
):
batch_id = tl.program_id(0)
hg_id = tl.program_id(1)
h_start = hg_id * HPB
h_off = h_start + tl.arange(0, HPB)
h_mask = h_off < nheads
q_s = tl.load(qo_indptr + batch_id)
kv_s = tl.load(kv_indptr + batch_id)
kv_e = tl.load(kv_indptr + batch_id + 1)
s_idx = tl.arange(0, BS)
dk_range = tl.arange(0, DK)
N_FULL_K: tl.constexpr = 576 // DK
REMAIN_K: tl.constexpr = 576 - N_FULL_K * DK
N_FULL_V: tl.constexpr = 512 // DK
combined_scale = sm_scale * kv_fp8_scale
q_pos = 0
while q_pos < qseqlen:
qtok = q_s + q_pos
mi = tl.full([HPB], float('-inf'), dtype=tl.float32)
li = tl.zeros([HPB], dtype=tl.float32)
a0 = tl.zeros([HPB, DK], dtype=tl.float32)
a1 = tl.zeros([HPB, DK], dtype=tl.float32)
a2 = tl.zeros([HPB, DK], dtype=tl.float32)
a3 = tl.zeros([HPB, DK], dtype=tl.float32)
for t_s in range(kv_s, kv_e, BS):
t_e = tl.minimum(t_s + BS, kv_e)
smask = s_idx < (t_e - t_s)
kv_ids = t_s + s_idx
scores = tl.zeros([HPB, BS], dtype=tl.float32)
for dc in range(N_FULL_K):
q_ptrs = Q + qtok * stride_qt + h_off[:, None] * stride_qh + dc * DK + dk_range[None, :]
q_chunk = tl.load(q_ptrs, mask=h_mask[:, None], other=0.0)
k_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + dc * DK + dk_range[None, :]
k_fp8 = tl.load(k_ptrs, mask=smask[:, None], other=0.0)
scores += tl.dot(q_chunk, tl.trans(k_fp8.to(tl.bfloat16)))
if REMAIN_K > 0:
r_dk = tl.arange(0, 64)
q_ptrs = Q + qtok * stride_qt + h_off[:, None] * stride_qh + N_FULL_K * DK + r_dk[None, :]
q_chunk = tl.load(q_ptrs, mask=h_mask[:, None], other=0.0)
k_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + N_FULL_K * DK + r_dk[None, :]
k_fp8 = tl.load(k_ptrs, mask=smask[:, None], other=0.0)
scores += tl.dot(q_chunk, tl.trans(k_fp8.to(tl.bfloat16)))
scores = scores * combined_scale
scores = tl.where(smask[None, :], scores, float('-inf'))
m_new = tl.max(scores, axis=1)
m_max = tl.maximum(mi, m_new)
exp_old = tl.math.exp2((mi - m_max) * LOG2E)
exp_s = tl.math.exp2((scores - m_max[:, None]) * LOG2E)
l_tile = tl.sum(exp_s, axis=1)
l_comb = li * exp_old + l_tile
safe_l = tl.where(l_comb > 0.0, l_comb, 1.0)
r = (li * exp_old / safe_l)[:, None]
a0 *= r; a1 *= r; a2 *= r; a3 *= r
probs = exp_s / safe_l[:, None]
pb = probs.to(tl.bfloat16)
for vc in range(N_FULL_V):
v_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + vc * DK + dk_range[None, :]
v_fp8 = tl.load(v_ptrs, mask=smask[:, None], other=0.0)
c = tl.dot(pb, v_fp8.to(tl.bfloat16))
if vc == 0: a0 += c
elif vc == 1: a1 += c
elif vc == 2: a2 += c
else: a3 += c
mi = m_max
li = l_comb
out_base = Out + qtok * stride_out_t
for vc in range(N_FULL_V):
if vc == 0: w = a0
elif vc == 1: w = a1
elif vc == 2: w = a2
else: w = a3
w_scaled = (w * kv_fp8_scale).to(tl.bfloat16)
tl.store(out_base + h_off[:, None] * stride_out_h + vc * DK + dk_range[None, :],
w_scaled, mask=h_mask[:, None])
q_pos += 1
@triton.jit
def _mla_splitk_small(
Q, KV_FP8,
pO, pM, pL,
qo_indptr, kv_indptr,
sm_scale, kv_fp8_scale,
stride_qt, stride_qh,
stride_kv,
stride_pO_s, stride_pO_t, stride_pO_h,
stride_pm_s, stride_pm_t,
nheads,
NSPLITS: tl.constexpr,
BS: tl.constexpr, HPB: tl.constexpr, QSEQLEN: tl.constexpr, DK: tl.constexpr,
):
pid0 = tl.program_id(0)
pid1 = tl.program_id(1)
batch_id = pid0 // NSPLITS
split_id = pid0 % NSPLITS
h_start = pid1 * HPB
h_off = h_start + tl.arange(0, HPB)
h_mask = h_off < nheads
q_s = tl.load(qo_indptr + batch_id)
kv_s = tl.load(kv_indptr + batch_id)
kv_e = tl.load(kv_indptr + batch_id + 1)
kv_len = kv_e - kv_s
chunk = tl.cdiv(kv_len, NSPLITS)
c_s = kv_s + split_id * chunk
c_e = tl.minimum(c_s + chunk, kv_e)
s_idx = tl.arange(0, BS)
dk_range = tl.arange(0, DK)
N_FULL_K: tl.constexpr = 576 // DK
REMAIN_K: tl.constexpr = 576 - N_FULL_K * DK
N_FULL_V: tl.constexpr = 512 // DK
for q_pos in range(QSEQLEN):
qtok = q_s + q_pos
mi = tl.full([HPB], float('-inf'), dtype=tl.float32)
li = tl.zeros([HPB], dtype=tl.float32)
a0 = tl.zeros([HPB, DK], dtype=tl.float32)
a1 = tl.zeros([HPB, DK], dtype=tl.float32)
a2 = tl.zeros([HPB, DK], dtype=tl.float32)
a3 = tl.zeros([HPB, DK], dtype=tl.float32)
for t_s in range(c_s, c_e, BS):
t_e = tl.minimum(t_s + BS, c_e)
smask = s_idx < (t_e - t_s)
kv_ids = t_s + s_idx
scores = tl.zeros([HPB, BS], dtype=tl.float32)
for dc in range(N_FULL_K):
q_ptrs = Q + qtok * stride_qt + h_off[:, None] * stride_qh + dc * DK + dk_range[None, :]
q_chunk = tl.load(q_ptrs, mask=h_mask[:, None], other=0.0)
k_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + dc * DK + dk_range[None, :]
k_fp8 = tl.load(k_ptrs, mask=smask[:, None], other=0.0)
scores += tl.dot(q_chunk, tl.trans(k_fp8.to(tl.bfloat16)))
if REMAIN_K > 0:
r_range = tl.arange(0, 64)
q_ptrs = Q + qtok * stride_qt + h_off[:, None] * stride_qh + N_FULL_K * DK + r_range[None, :]
q_chunk = tl.load(q_ptrs, mask=h_mask[:, None], other=0.0)
k_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + N_FULL_K * DK + r_range[None, :]
k_fp8 = tl.load(k_ptrs, mask=smask[:, None], other=0.0)
scores += tl.dot(q_chunk, tl.trans(k_fp8.to(tl.bfloat16)))
scores = scores * sm_scale
scores = tl.where(smask[None, :], scores, float('-inf'))
m_new = tl.max(scores, axis=1)
m_max = tl.maximum(mi, m_new)
exp_old = tl.math.exp2((mi - m_max) * LOG2E)
exp_s = tl.math.exp2((scores - m_max[:, None]) * LOG2E)
l_tile = tl.sum(exp_s, axis=1)
l_comb = li * exp_old + l_tile
safe_l = tl.where(l_comb > 0.0, l_comb, 1.0)
r = (li * exp_old / safe_l)[:, None]
a0 *= r; a1 *= r; a2 *= r; a3 *= r
probs = exp_s / safe_l[:, None]
pb = probs.to(tl.bfloat16)
for vc in range(N_FULL_V):
v_ptrs = KV_FP8 + kv_ids[:, None] * stride_kv + vc * DK + dk_range[None, :]
v_fp8 = tl.load(v_ptrs, mask=smask[:, None], other=0.0)
c = tl.dot(pb, v_fp8.to(tl.bfloat16))
if vc == 0:
a0 += c
elif vc == 1:
a1 += c
elif vc == 2:
a2 += c
else:
a3 += c
mi = m_max
li = l_comb
a0 *= kv_fp8_scale; a1 *= kv_fp8_scale; a2 *= kv_fp8_scale; a3 *= kv_fp8_scale
po_base = pO + split_id * stride_pO_s + qtok * stride_pO_t
for vc in range(N_FULL_V):
if vc == 0:
w = a0
elif vc == 1:
w = a1
elif vc == 2:
w = a2
else:
w = a3
tl.store(po_base + h_off[:, None] * stride_pO_h + vc * DK + dk_range[None, :],
w, mask=h_mask[:, None])
pm_off = pM + split_id * stride_pm_s + qtok * stride_pm_t + h_off
pl_off = pL + split_id * stride_pm_s + qtok * stride_pm_t + h_off
tl.store(pm_off, mi, mask=h_mask)
tl.store(pl_off, li, mask=h_mask)
@triton.jit
def _mla_reduce_small(
pO, pM, pL, Out,
stride_pO_s, stride_pm_s,
NSPLITS: tl.constexpr,
VDIM: tl.constexpr,
BLK: tl.constexpr,
):
th = tl.program_id(0)
gm = tl.full([], float('-inf'), dtype=tl.float32)
for s in range(NSPLITS):
gm = tl.maximum(gm, tl.load(pM + s * stride_pm_s + th))
gl = tl.zeros([], dtype=tl.float32)
for s in range(NSPLITS):
pm = tl.load(pM + s * stride_pm_s + th)
pl = tl.load(pL + s * stride_pm_s + th)
gl += pl * tl.math.exp2((pm - gm) * LOG2E)
safe_gl = tl.where(gl > 0.0, gl, 1.0)
d = tl.arange(0, BLK)
for d_start in range(0, VDIM, BLK):
acc = tl.zeros([BLK], dtype=tl.float32)
for s in range(NSPLITS):
pm = tl.load(pM + s * stride_pm_s + th)
pl = tl.load(pL + s * stride_pm_s + th)
w = pl * tl.math.exp2((pm - gm) * LOG2E)
acc += tl.load(pO + s * stride_pO_s + th * VDIM + d_start + d) * w
tl.store(Out + th * VDIM + d_start + d, (acc / safe_gl).to(tl.bfloat16))
# ═══════════════════════════════════════════════════════════════════════
# MXFP4 kernels — ~1.9x bandwidth savings for large configs
# ═══════════════════════════════════════════════════════════════════════
@triton.jit
def _fp4e2m1_lookup(u):
"""Convert unsigned 3-bit fp4 magnitude (0-7) to float32."""
return tl.where(u < 4,
tl.where(u < 2,
tl.where(u == 0, 0.0, 0.5),
tl.where(u == 2, 1.0, 1.5)),
tl.where(u < 6,
tl.where(u == 4, 2.0, 3.0),
tl.where(u == 6, 4.0, 6.0)))
@triton.jit
def _mla_hybrid_splitk(
Q_FP8, KV_PACKED, KV_SCALE,
pO, pM, pL,
qo_indptr, kv_indptr,
sm_scale,
stride_qt, stride_qh,
stride_kvp, stride_ks,
stride_pO_s, stride_pO_t, stride_pO_h,
stride_pm_s, stride_pm_t,
nheads,
NSPLITS: tl.constexpr,
BS: tl.constexpr, HPB: tl.constexpr, DK: tl.constexpr,
):
N_K_FULL: tl.constexpr = 576 // DK
K_REM: tl.constexpr = 576 - N_K_FULL * DK
K_REM_PACKED: tl.constexpr = K_REM // 2
K_REM_NSB: tl.constexpr = K_REM // 32
N_V: tl.constexpr = 512 // DK
PKD: tl.constexpr = DK // 2
NSB: tl.constexpr = DK // 32
HALF_DK: tl.constexpr = DK // 2
pid0 = tl.program_id(0)
hg_id = tl.program_id(1)
batch_id = pid0 // NSPLITS
split_id = pid0 % NSPLITS
h_start = hg_id * HPB
h_off = h_start + tl.arange(0, HPB)
h_mask = h_off < nheads
q_s = tl.load(qo_indptr + batch_id)
kv_s = tl.load(kv_indptr + batch_id)
kv_e = tl.load(kv_indptr + batch_id + 1)
kv_len = kv_e - kv_s
chunk_size = tl.cdiv(kv_len, NSPLITS)
c_s = kv_s + split_id * chunk_size
c_e = tl.minimum(c_s + chunk_size, kv_e)
s_idx = tl.arange(0, BS)
dk_range = tl.arange(0, DK)
pk_range = tl.arange(0, PKD)
sb_range = tl.arange(0, NSB)
half_range = tl.arange(0, HALF_DK)
scale_map = pk_range // 16
qtok = q_s
mi = tl.full([HPB], float('-inf'), dtype=tl.float32)
li = tl.zeros([HPB], dtype=tl.float32)
ae0 = tl.zeros([HPB, HALF_DK], dtype=tl.float32)
ao0 = tl.zeros([HPB, HALF_DK], dtype=tl.float32)
ae1 = tl.zeros([HPB, HALF_DK], dtype=tl.float32)
ao1 = tl.zeros([HPB, HALF_DK], dtype=tl.float32)
ae2 = tl.zeros([HPB, HALF_DK], dtype=tl.float32)
ao2 = tl.zeros([HPB, HALF_DK], dtype=tl.float32)
ae3 = tl.zeros([HPB, HALF_DK], dtype=tl.float32)
ao3 = tl.zeros([HPB, HALF_DK], dtype=tl.float32)
for t_s in range(c_s, c_e, BS):
t_e = tl.minimum(t_s + BS, c_e)
smask = s_idx < (t_e - t_s)
kv_ids = t_s + s_idx
scores = tl.zeros([HPB, BS], dtype=tl.float32)
for dc in range(N_K_FULL):
q_ptrs = Q_FP8 + qtok * stride_qt + h_off[:, None] * stride_qh + dc * DK + dk_range[None, :]
q_tile = tl.load(q_ptrs, mask=h_mask[:, None], other=0.0)
k_ptrs = KV_PACKED + kv_ids[:, None] * stride_kvp + dc * PKD + pk_range[None, :]
k_packed = tl.load(k_ptrs, mask=smask[:, None], other=0)
k_t = tl.trans(k_packed)
ks_ptrs = KV_SCALE + kv_ids[:, None] * stride_ks + dc * NSB + sb_range[None, :]
k_scale = tl.load(ks_ptrs, mask=smask[:, None], other=127)
scores += tl.dot_scaled(q_tile, None, "e4m3", k_t, k_scale, "e2m1")
if K_REM > 0:
rem_range = tl.arange(0, K_REM)
rem_pk_range = tl.arange(0, K_REM_PACKED)
rem_sb_range = tl.arange(0, K_REM_NSB)
qr_ptrs = Q_FP8 + qtok * stride_qt + h_off[:, None] * stride_qh + N_K_FULL * DK + rem_range[None, :]
qr = tl.load(qr_ptrs, mask=h_mask[:, None], other=0.0)
kr_ptrs = KV_PACKED + kv_ids[:, None] * stride_kvp + N_K_FULL * PKD + rem_pk_range[None, :]
kr = tl.load(kr_ptrs, mask=smask[:, None], other=0)
kr_t = tl.trans(kr)
ksr_ptrs = KV_SCALE + kv_ids[:, None] * stride_ks + N_K_FULL * NSB + rem_sb_range[None, :]
kr_scale = tl.load(ksr_ptrs, mask=smask[:, None], other=127)
scores += tl.dot_scaled(qr, None, "e4m3", kr_t, kr_scale, "e2m1")
scores = scores * sm_scale
scores = tl.where(smask[None, :], scores, float('-inf'))
m_new = tl.max(scores, axis=1)
m_max = tl.maximum(mi, m_new)
exp_old = tl.math.exp2((mi - m_max) * LOG2E)
exp_s = tl.math.exp2((scores - m_max[:, None]) * LOG2E)
l_tile = tl.sum(exp_s, axis=1)
l_comb = li * exp_old + l_tile
safe_l = tl.where(l_comb > 0.0, l_comb, 1.0)
r = (li * exp_old / safe_l)[:, None]
ae0 *= r; ao0 *= r; ae1 *= r; ao1 *= r
ae2 *= r; ao2 *= r; ae3 *= r; ao3 *= r
probs = exp_s / safe_l[:, None]
pb = probs.to(tl.bfloat16)
for vc in range(N_V):
vp_ptrs = KV_PACKED + kv_ids[:, None] * stride_kvp + vc * PKD + pk_range[None, :]
vp = tl.load(vp_ptrs, mask=smask[:, None], other=0).to(tl.uint8)
sc_col = vc * NSB + scale_map
vs_ptrs = KV_SCALE + kv_ids[:, None] * stride_ks + sc_col[None, :]
vs = tl.math.exp2(tl.load(vs_ptrs, mask=smask[:, None], other=127).to(tl.float32) - 127.0)
vl = _fp4e2m1_lookup(vp & 7) * (1.0 - ((vp >> 3) & 1).to(tl.float32) * 2.0) * vs
vh = _fp4e2m1_lookup((vp >> 4) & 7) * (1.0 - ((vp >> 7) & 1).to(tl.float32) * 2.0) * vs
ce = tl.dot(pb, vl.to(tl.bfloat16))
co = tl.dot(pb, vh.to(tl.bfloat16))
if vc == 0: ae0 += ce; ao0 += co
elif vc == 1: ae1 += ce; ao1 += co
elif vc == 2: ae2 += ce; ao2 += co
else: ae3 += ce; ao3 += co
mi = m_max
li = l_comb
po_base = pO + split_id * stride_pO_s + qtok * stride_pO_t
for vc in range(N_V):
if vc == 0: we, wo = ae0, ao0
elif vc == 1: we, wo = ae1, ao1
elif vc == 2: we, wo = ae2, ao2
else: we, wo = ae3, ao3
even_offs = vc * DK + 2 * half_range
odd_offs = even_offs + 1
tl.store(po_base + h_off[:, None] * stride_pO_h + even_offs[None, :], we, mask=h_mask[:, None])
tl.store(po_base + h_off[:, None] * stride_pO_h + odd_offs[None, :], wo, mask=h_mask[:, None])
pm_off = pM + split_id * stride_pm_s + qtok * stride_pm_t + h_off
pl_off = pL + split_id * stride_pm_s + qtok * stride_pm_t + h_off
tl.store(pm_off, mi, mask=h_mask)
tl.store(pl_off, li, mask=h_mask)
@triton.jit
def _mla_mxfp4_reduce(
pO, pM, pL, Out,
stride_pO_s, stride_pO_t, stride_pO_h,
stride_pm_s, stride_pm_t,
stride_out_t, stride_out_h,
nheads,
NSPLITS: tl.constexpr,
HPB: tl.constexpr, DK: tl.constexpr,
):
qtok = tl.program_id(0)
hg_id = tl.program_id(1)
h_start = hg_id * HPB
h_off = h_start + tl.arange(0, HPB)
h_mask = h_off < nheads
dk_range = tl.arange(0, DK)
N_FULL_V: tl.constexpr = 512 // DK
gm = tl.full([HPB], float('-inf'), dtype=tl.float32)
gl = tl.zeros([HPB], dtype=tl.float32)
ra0 = tl.zeros([HPB, DK], dtype=tl.float32)
ra1 = tl.zeros([HPB, DK], dtype=tl.float32)
ra2 = tl.zeros([HPB, DK], dtype=tl.float32)
ra3 = tl.zeros([HPB, DK], dtype=tl.float32)
for s in range(NSPLITS):
pm_val = tl.load(pM + s * stride_pm_s + qtok * stride_pm_t + h_off,
mask=h_mask, other=float('-inf'))
pl_val = tl.load(pL + s * stride_pm_s + qtok * stride_pm_t + h_off,
mask=h_mask, other=0.0)
m_new = tl.maximum(gm, pm_val)
exp_old = tl.math.exp2((gm - m_new) * LOG2E)
exp_new = tl.math.exp2((pm_val - m_new) * LOG2E)
gl_new = gl * exp_old + pl_val * exp_new
safe_gl = tl.where(gl_new > 0.0, gl_new, 1.0)
r = (gl * exp_old / safe_gl)[:, None]
ra0 *= r; ra1 *= r; ra2 *= r; ra3 *= r
f = (pl_val * exp_new / safe_gl)[:, None]
po_base = pO + s * stride_pO_s + qtok * stride_pO_t + h_off[:, None] * stride_pO_h
for vc in range(N_FULL_V):
partial = tl.load(po_base + vc * DK + dk_range[None, :],
mask=h_mask[:, None], other=0.0)
c = partial * f
if vc == 0: ra0 += c
elif vc == 1: ra1 += c
elif vc == 2: ra2 += c
else: ra3 += c
gm = m_new
gl = gl_new
out_base = Out + qtok * stride_out_t + h_off[:, None] * stride_out_h
for vc in range(N_FULL_V):
if vc == 0: w = ra0
elif vc == 1: w = ra1
elif vc == 2: w = ra2
else: w = ra3
tl.store(out_base + vc * DK + dk_range[None, :],
w.to(tl.bfloat16), mask=h_mask[:, None])
# ═══════════════════════════════════════════════════════════════════════
# Warmup
# ═══════════════════════════════════════════════════════════════════════
_warmed = False
def _do_warmup():
global _warmed
if _warmed:
return
_warmed = True
BS_W = 128; DK_W = 128
qi_w = torch.tensor([0, 1], dtype=torch.int32, device="cuda")
ki_w = torch.tensor([0, BS_W + 1], dtype=torch.int32, device="cuda")
kv_fp8_w = torch.zeros((BS_W + 1, QK_DIM), dtype=FP8, device="cuda")
kv_mx_w = torch.zeros((BS_W + 1, QK_DIM // 2), dtype=torch.uint8, device="cuda")
kv_ms_w = torch.full((BS_W + 1, 24), 127, dtype=torch.uint8, device="cuda")
for HPB_W in [16, 64]:
for n_hg_v in [1, 2]:
nh_v = n_hg_v * HPB_W
q_w = torch.zeros((1, nh_v, QK_DIM), dtype=torch.bfloat16, device="cuda")
qf_w = torch.zeros((1, nh_v, QK_DIM), dtype=FP8, device="cuda")
o_w = torch.zeros((1, nh_v, V_DIM), dtype=torch.bfloat16, device="cuda")
_mla_nosplit[(1, n_hg_v)](
q_w, kv_fp8_w, o_w,
qi_w, ki_w,
1.0, 1.0,
q_w.stride(0), q_w.stride(1),
kv_fp8_w.stride(0),
o_w.stride(0), o_w.stride(1),
nh_v, 1,
BS=BS_W, HPB=HPB_W, DK=DK_W,
num_warps=4, num_stages=2,
)
for NS in [2, 4, 8, 16, 32, 64]:
pO_w = torch.zeros((NS, 1, nh_v, V_DIM), dtype=torch.float32, device="cuda")
pM_w = torch.zeros((NS, 1, nh_v), dtype=torch.float32, device="cuda")
pL_w = torch.zeros((NS, 1, nh_v), dtype=torch.float32, device="cuda")
ctr_w = torch.zeros(n_hg_v, dtype=torch.int32, device="cuda")
_mla_fused[(1, NS, n_hg_v)](
q_w, kv_fp8_w,
pO_w, pM_w, pL_w, o_w,
qi_w, ki_w, ctr_w,
1.0, 1.0,
q_w.stride(0), q_w.stride(1),
kv_fp8_w.stride(0),
pO_w.stride(0), pO_w.stride(1), pO_w.stride(2),
pM_w.stride(0), pM_w.stride(1),
o_w.stride(0), o_w.stride(1),
NS, nh_v, 1, n_hg_v,
NSPLITS=NS, BS=BS_W, HPB=HPB_W, DK=DK_W,
num_warps=4, num_stages=2,
)
if HPB_W == 16:
for NS in [4, 8]:
pO_w = torch.zeros((NS, 1, nh_v, V_DIM), dtype=torch.float32, device="cuda")
pM_w = torch.zeros((NS, 1, nh_v), dtype=torch.float32, device="cuda")
pL_w = torch.zeros((NS, 1, nh_v), dtype=torch.float32, device="cuda")
_mla_splitk_small[(NS, n_hg_v)](
q_w, kv_fp8_w,
pO_w, pM_w, pL_w,
qi_w, ki_w,
1.0, 1.0,
q_w.stride(0), q_w.stride(1),
kv_fp8_w.stride(0),
pO_w.stride(0), pO_w.stride(1), pO_w.stride(2),
pM_w.stride(0), pM_w.stride(1),
nh_v,
NSPLITS=NS, BS=BS_W, HPB=HPB_W, QSEQLEN=1, DK=DK_W,
num_warps=4, num_stages=2,
)
pO_flat_w = pO_w.view(NS, -1)
pM_flat_w = pM_w.view(NS, -1)
pL_flat_w = pL_w.view(NS, -1)
_mla_reduce_small[(nh_v,)](
pO_flat_w, pM_flat_w, pL_flat_w, o_w.view(-1),
pO_flat_w.stride(0), pM_flat_w.stride(0),
NSPLITS=NS, VDIM=V_DIM, BLK=128,
num_warps=2, num_stages=1,
)
for NS in [4, 8, 16, 32]:
pO_w = torch.zeros((NS, 1, nh_v, V_DIM), dtype=torch.float32, device="cuda")
pM_w = torch.zeros((NS, 1, nh_v), dtype=torch.float32, device="cuda")
pL_w = torch.zeros((NS, 1, nh_v), dtype=torch.float32, device="cuda")
qf_w.zero_()
_mla_hybrid_splitk[(NS, n_hg_v)](
qf_w, kv_mx_w, kv_ms_w,
pO_w, pM_w, pL_w,
qi_w, ki_w,
1.0,
qf_w.stride(0), qf_w.stride(1),
kv_mx_w.stride(0), kv_ms_w.stride(0),
pO_w.stride(0), pO_w.stride(1), pO_w.stride(2),
pM_w.stride(0), pM_w.stride(1),
nh_v,
NSPLITS=NS, BS=BS_W, HPB=HPB_W, DK=DK_W,
num_warps=4, num_stages=2,
)
_mla_mxfp4_reduce[(1, n_hg_v)](
pO_w, pM_w, pL_w, o_w,
pO_w.stride(0), pO_w.stride(1), pO_w.stride(2),
pM_w.stride(0), pM_w.stride(1),
o_w.stride(0), o_w.stride(1),
nh_v,
NSPLITS=NS, HPB=HPB_W, DK=DK_W,
num_warps=2, num_stages=1,
)
torch.cuda.synchronize()
torch.cuda.empty_cache()
# ═══════════════════════════════════════════════════════════════════════
# Aiter FP8 — lean dispatch with pre-computed args
# ═══════════════════════════════════════════════════════════════════════
_NS_TABLE = {
(4, 1024): 16, (4, 8192): 32,
(32, 1024): 16, (32, 8192): 32,
(64, 1024): 16, (64, 8192): 64,
(256, 1024): 64, (256, 8192): 64,
}
class _AiterEntry:
__slots__ = ['qf_buf', 'qf_view', 'o', 'kvb4', 'kvi', 'klp',
'qs', 'kvs', 'meta', 'ns', 'sms', 'nkv', 'qsl',
'wk', 'logits', 'attn_lse', 'q_folded', 'o_folded',
'use_direct', 'direct_safe', 'q_holder', 'graph',
'qo_indptr_own', 'kv_indptr_own']
_aiter_table = {}
def _aiter_setup(q, kv_data, qo_indptr, kv_indptr, config, key):
FP8_A = aiter_dtypes.fp8
nh = config["num_heads"]; nkv = config["num_kv_heads"]
dqk = config["qk_head_dim"]; dv = config["v_head_dim"]
bs = config["batch_size"]; qsl = config["q_seq_len"]
kvsl = config["kv_seq_len"]; sms = config["sm_scale"]
qf_buf = torch.empty(q.shape, dtype=FP8_A, device="cuda")
qs = torch.ones(1, dtype=torch.float32, device="cuda")
kvb, kvs = kv_data["fp8"]
tkv = bs * kvsl
kvi = torch.arange(tkv, dtype=torch.int32, device="cuda")
klp = torch.full((bs,), kvsl, dtype=torch.int32, device="cuda")
kvb4 = kvb.view(kvb.shape[0], 1, nkv, kvb.shape[-1])
qo_own = torch.arange(0, bs + 1, dtype=torch.int32, device="cuda") * qsl
kv_own = torch.arange(0, (bs + 1) * kvsl, kvsl, dtype=torch.int32, device="cuda")
ns = 64 if qsl == 1 else _NS_TABLE.get((bs, kvsl), 64)
ibm = True
info = get_mla_metadata_info_v1(
bs, qsl, nh, qf_buf.dtype, kvb.dtype, is_sparse=False, fast_mode=False,
num_kv_splits=ns, intra_batch_mode=ibm)
wk = [torch.empty(s, dtype=t, device="cuda") for s, t in info]
get_mla_metadata_v1(
qo_indptr, kv_indptr, klp, nh // nkv, nkv, True,
wk[0], wk[2], wk[1], wk[3], wk[4], wk[5],
page_size=1, kv_granularity=16,
max_seqlen_qo=qsl, uni_seqlen_qo=qsl, fast_mode=False,
max_split_per_batch=ns, intra_batch_mode=ibm,
dtype_q=qf_buf.dtype, dtype_kv=kvb.dtype)
total_q = q.shape[0]
e = _AiterEntry()
e.qf_buf = qf_buf
e.qf_view = qf_buf.view(-1, nh, dqk)
e.o = torch.empty((total_q, nh, dv), dtype=torch.bfloat16, device="cuda")
e.kvb4 = kvb4
e.kvi = kvi
e.klp = klp
e.qs = qs
e.kvs = kvs
e.meta = dict(
work_meta_data=wk[0], work_indptr=wk[1], work_info_set=wk[2],
reduce_indptr=wk[3], reduce_final_map=wk[4], reduce_partial_map=wk[5])
e.ns = ns
e.sms = sms
e.nkv = nkv
e.qsl = qsl
e.wk = wk
e.qo_indptr_own = qo_own
e.kv_indptr_own = kv_own
e.use_direct = False
e.direct_safe = False
e.graph = None
e.q_holder = None
if _direct_available:
try:
if qsl == 1:
fold = max(nh // 16, 1)
nh_f = 16 if nh >= 16 else nh
total_s_f = total_q * fold
e.q_folded = e.qf_view.view(total_s_f, nh_f, dqk)
e.o_folded = e.o.view(total_s_f, nh_f, dv)
else:
raise RuntimeError(f"direct path unsupported for nh={nh}, qsl={qsl}")
rpm_sz = wk[5].size(0)
e.logits = torch.empty((rpm_sz * qsl, 1, nh_f, dv),
dtype=torch.float32, device="cuda")
e.attn_lse = torch.empty((rpm_sz * qsl, 1, nh_f, 1),
dtype=torch.float32, device="cuda")
e.q_holder = torch.empty_like(q)
e.use_direct = True
e.direct_safe = True
# Warm up the direct path to trigger JIT/compilation
e.q_holder.copy_(q)
e.qf_buf.copy_(e.q_holder)
_stage1_fn(
e.q_folded, kvb4, qo_own, kv_own, kvi, klp,
None, wk[0], wk[1], wk[2],
qsl, 1, nkv, sms,
e.logits, e.attn_lse, e.o_folded,
qs, kvs)
_reduce_fn(
e.logits, e.attn_lse, wk[3], wk[4], wk[5],
qsl, e.o_folded, None)
torch.cuda.synchronize()
except Exception as ex:
print(f"[MLA] Direct/graph setup failed bs={bs},kvsl={kvsl}: {ex}", file=sys.stderr)
e.use_direct = False
e.graph = None
_aiter_table[key] = e
return e
def _aiter_refresh_inputs(e, kv_data):
kvb, kvs = kv_data["fp8"]
e.kvb4 = kvb.view(kvb.shape[0], 1, e.nkv, kvb.shape[-1])
e.kvs = kvs
# ═══════════════════════════════════════════════════════════════════════
# Host dispatch
# ═══════════════════════════════════════════════════════════════════════
_prev_buf_key = None
_prev_bufs = None
_prev_out_key = None
_prev_out = None
_counter_cache = {}
_gen_counter = {}
_qo_merged_cache = {}
_config_cache = {}
_mxfp4_config_cache = {}
_q_fp8_cache = {}
_kv_fp8_view_cache = {}
_kv_fp8_scale_cache = {}
_SMALL_TRITON_NS = {
4: 8,
32: 8,
64: 4,
}
def _get_bufs(nsplits, total_q, nh):
global _prev_buf_key, _prev_bufs
key = (nsplits, total_q, nh)
if key != _prev_buf_key:
pO = torch.empty((nsplits, total_q, nh, V_DIM), dtype=torch.float32, device="cuda")
pM = torch.empty((nsplits, total_q, nh), dtype=torch.float32, device="cuda")
pL = torch.empty((nsplits, total_q, nh), dtype=torch.float32, device="cuda")
_prev_buf_key = key
_prev_bufs = (pO, pM, pL)
return _prev_bufs
def _get_output(total_q, nh):
global _prev_out_key, _prev_out
key = (total_q, nh)
if key != _prev_out_key:
_prev_out_key = key
_prev_out = torch.empty((total_q, nh, V_DIM), dtype=torch.bfloat16, device="cuda")
return _prev_out
def _get_counter(bs, n_hg):
key = (bs, n_hg)
if key not in _counter_cache:
_counter_cache[key] = torch.zeros(bs * n_hg, dtype=torch.int32, device="cuda")
_gen_counter[key] = 0
return _counter_cache[key]
def _get_qo_merged(bs):
if bs not in _qo_merged_cache:
_qo_merged_cache[bs] = torch.arange(bs + 1, dtype=torch.int32, device="cuda")
return _qo_merged_cache[bs]
def _get_q_fp8(shape):
key = tuple(shape)
if key not in _q_fp8_cache:
_q_fp8_cache[key] = torch.empty(shape, dtype=FP8, device="cuda")
return _q_fp8_cache[key]
def _get_kv_fp8_view(kvb_fp8):
key = kvb_fp8.data_ptr()
view = _kv_fp8_view_cache.get(key)
if view is None:
view = kvb_fp8.reshape(-1, QK_DIM)
_kv_fp8_view_cache[key] = view
return view
def _get_kv_fp8_scale(kvs_fp8):
key = kvs_fp8.data_ptr()
scale = _kv_fp8_scale_cache.get(key)
if scale is None:
scale = kvs_fp8.float().item()
_kv_fp8_scale_cache[key] = scale
return scale
def _get_uint8_view(tensor):
key = tensor.data_ptr()
view = _config_cache.get(("view", key))
if view is None:
flat = tensor.reshape(-1, tensor.shape[-1])
view = flat.view(torch.uint8).reshape(flat.shape[0], flat.shape[1])
_config_cache[("view", key)] = view
return view
def _should_use_mxfp4(config):
bs = config["batch_size"]
qsl = config["q_seq_len"]
kvsl = config["kv_seq_len"]
return qsl > 1 and kvsl >= 8192 and bs >= 32
def _should_use_small_triton(config):
return (config["q_seq_len"] == 1 and
config["kv_seq_len"] == 1024 and
config["batch_size"] in _SMALL_TRITON_NS)
def _run_mxfp4(q, kv_data, qo_indptr, kv_indptr, config):
nh = config["num_heads"]
bs = config["batch_size"]
qsl = config["q_seq_len"]
kvsl = config["kv_seq_len"]
sms = config["sm_scale"]
total_q = q.shape[0]
kv_mxfp4_raw, kv_scale_raw = kv_data["mxfp4"]
kv_mxfp4 = _get_uint8_view(kv_mxfp4_raw)
kv_scale = _get_uint8_view(kv_scale_raw)
if qsl > 1:
nh_eff = qsl * nh
total_q_eff = bs
q_for_kernel = q.reshape(bs, nh_eff, QK_DIM)
qo_for_kernel = _get_qo_merged(bs)
else:
nh_eff = nh
total_q_eff = total_q
q_for_kernel = q
qo_for_kernel = qo_indptr
HPB = 64 if nh_eff >= 64 else 16
DK = 128
BS = 128
n_hg = triton.cdiv(nh_eff, HPB)
q_fp8 = _get_q_fp8(q_for_kernel.shape)
q_fp8.copy_(q_for_kernel)
cfg_key = (bs, nh_eff, kvsl, HPB)
if cfg_key not in _mxfp4_config_cache:
grid_per_split = bs * n_hg
max_splits_kv = max(1, kvsl // BS)
target_wgs = 304 * 2
raw_splits = max(1, target_wgs // max(grid_per_split, 1))
raw_splits = min(raw_splits, max_splits_kv)
ns = 1
for candidate in [1, 2, 4, 8, 16, 32, 64]:
if candidate <= raw_splits:
ns = candidate
_mxfp4_config_cache[cfg_key] = ns
nsplits = _mxfp4_config_cache[cfg_key]
pO, pM, pL = _get_bufs(nsplits, total_q_eff, nh_eff)
output = _get_output(total_q_eff, nh_eff)
_mla_hybrid_splitk[(bs * nsplits, n_hg)](
q_fp8, kv_mxfp4, kv_scale,
pO, pM, pL,
qo_for_kernel, kv_indptr,
sms,
q_fp8.stride(0), q_fp8.stride(1),
kv_mxfp4.stride(0), kv_scale.stride(0),
pO.stride(0), pO.stride(1), pO.stride(2),
pM.stride(0), pM.stride(1),
nh_eff,
NSPLITS=nsplits, BS=BS, HPB=HPB, DK=DK,
num_warps=4, num_stages=2,
)
_mla_mxfp4_reduce[(total_q_eff, n_hg)](
pO, pM, pL, output,
pO.stride(0), pO.stride(1), pO.stride(2),
pM.stride(0), pM.stride(1),
output.stride(0), output.stride(1),
nh_eff,
NSPLITS=nsplits, HPB=HPB, DK=DK,
num_warps=2, num_stages=1,
)
if qsl > 1:
return output.reshape(bs, qsl, nh, V_DIM).reshape(total_q, nh, V_DIM)
return output
def _run_small_triton(q, kv_data, qo_indptr, kv_indptr, config):
nh = config["num_heads"]
bs = config["batch_size"]
sms = config["sm_scale"]
total_q = q.shape[0]
kvb_fp8, kvs_fp8 = kv_data["fp8"]
kv_fp8 = _get_kv_fp8_view(kvb_fp8)
kv_fp8_scale_val = _get_kv_fp8_scale(kvs_fp8)
effective_sms = sms * kv_fp8_scale_val
HPB = 16
DK = 128
BS = 128
n_hg = triton.cdiv(nh, HPB)
nsplits = _SMALL_TRITON_NS[bs]
output = _get_output(total_q, nh)
if nsplits <= 1:
_mla_nosplit[(bs, n_hg)](
q, kv_fp8, output,
qo_indptr, kv_indptr,
sms, kv_fp8_scale_val,
q.stride(0), q.stride(1),
kv_fp8.stride(0),
output.stride(0), output.stride(1),
nh, 1,
BS=BS, HPB=HPB, DK=DK,
num_warps=4, num_stages=2,
)
return output
pO, pM, pL = _get_bufs(nsplits, total_q, nh)
_mla_splitk_small[(bs * nsplits, n_hg)](
q, kv_fp8,
pO, pM, pL,
qo_indptr, kv_indptr,
effective_sms, kv_fp8_scale_val,
q.stride(0), q.stride(1),
kv_fp8.stride(0),
pO.stride(0), pO.stride(1), pO.stride(2),
pM.stride(0), pM.stride(1),
nh,
NSPLITS=nsplits, BS=BS, HPB=HPB, QSEQLEN=1, DK=DK,
num_warps=4, num_stages=2,
)
pO_flat = pO.view(nsplits, -1)
pM_flat = pM.view(nsplits, -1)
pL_flat = pL.view(nsplits, -1)
_mla_reduce_small[(total_q * nh,)](
pO_flat, pM_flat, pL_flat, output.view(-1),
pO_flat.stride(0), pM_flat.stride(0),
NSPLITS=nsplits, VDIM=V_DIM, BLK=128,
num_warps=2, num_stages=1,
)
return output
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
if _should_use_mxfp4(config):
try:
_do_warmup()
return _run_mxfp4(q, kv_data, qo_indptr, kv_indptr, config)
except Exception as ex:
print(f"[MLA] MXFP4 path failed, fallback: {ex}", file=sys.stderr)
traceback.print_exc(file=sys.stderr)
if _should_use_small_triton(config):
try:
_do_warmup()
return _run_small_triton(q, kv_data, qo_indptr, kv_indptr, config)
except Exception as ex:
print(f"[MLA] Small Triton path failed, fallback: {ex}", file=sys.stderr)
traceback.print_exc(file=sys.stderr)
if _aiter_available:
key = (
config["batch_size"],
config["kv_seq_len"],
config["num_heads"],
config["q_seq_len"],
)
e = _aiter_table.get(key)
if e is None:
e = _aiter_setup(q, kv_data, qo_indptr, kv_indptr, config, key)
_aiter_refresh_inputs(e, kv_data)
if e.graph is not None:
e.q_holder.copy_(q)
e.graph.replay()
return e.o
if e.use_direct:
e.qf_buf.copy_(q)
try:
_stage1_fn(
e.q_folded, e.kvb4, e.qo_indptr_own, e.kv_indptr_own,
e.kvi, e.klp,
None, e.wk[0], e.wk[1], e.wk[2],
e.qsl, 1, e.nkv, e.sms,
e.logits, e.attn_lse, e.o_folded,
e.qs, e.kvs)
_reduce_fn(
e.logits, e.attn_lse, e.wk[3], e.wk[4], e.wk[5],
e.qsl, e.o_folded, None)
return e.o
except Exception as ex:
print(f"[MLA] Direct call failed, fallback: {ex}", file=sys.stderr)
e.use_direct = False
e.qf_buf.copy_(q)
mla_decode_fwd(
e.qf_view, e.kvb4, e.o, qo_indptr, kv_indptr, e.kvi, e.klp,
e.qsl, page_size=1, nhead_kv=e.nkv, sm_scale=e.sms, logit_cap=0.0,
num_kv_splits=e.ns, q_scale=e.qs, kv_scale=e.kvs,
intra_batch_mode=True, **e.meta)
return e.o
_do_warmup()
return _run_triton(q, kv_data, qo_indptr, kv_indptr, config)
def _run_triton(q, kv_data, qo_indptr, kv_indptr, config):
nh = config["num_heads"]
bs = config["batch_size"]
qsl = config["q_seq_len"]
kvsl = config["kv_seq_len"]
sms = config["sm_scale"]
total_q = q.shape[0]
kvb_fp8, kvs_fp8 = kv_data["fp8"]
kv_fp8 = _get_kv_fp8_view(kvb_fp8)
kv_fp8_scale_val = _get_kv_fp8_scale(kvs_fp8)
if qsl > 1:
nh_eff = qsl * nh
total_q_eff = bs
q_for_kernel = q.reshape(bs, nh_eff, QK_DIM)
qo_for_kernel = _get_qo_merged(bs)
qsl_eff = 1
else:
nh_eff = nh
total_q_eff = total_q
q_for_kernel = q
qo_for_kernel = qo_indptr
qsl_eff = 1
if qsl > 1 and nh_eff >= 64:
HPB = 64
else:
HPB = 16
DK = 128; BS = 128
num_warps = 4; num_stages = 2
n_hg = triton.cdiv(nh_eff, HPB)
config_key = (bs, nh_eff, kvsl, HPB)
if config_key not in _config_cache:
grid_per_split = bs * n_hg
max_splits_kv = max(1, kvsl // BS)
target_wgs = 304 * 2
raw_splits = max(1, target_wgs // max(grid_per_split, 1))
raw_splits = min(raw_splits, max_splits_kv)
NSPLITS = 1
for s in [1, 2, 4, 8, 16, 32, 64]:
if s <= raw_splits:
NSPLITS = s
_config_cache[config_key] = NSPLITS
NSPLITS = _config_cache[config_key]
output = _get_output(total_q_eff, nh_eff)
if NSPLITS <= 1:
_mla_nosplit[(bs, n_hg)](
q_for_kernel, kv_fp8, output,
qo_for_kernel, kv_indptr,
sms, kv_fp8_scale_val,
q_for_kernel.stride(0), q_for_kernel.stride(1),
kv_fp8.stride(0),
output.stride(0), output.stride(1),
nh_eff, qsl_eff,
BS=BS, HPB=HPB, DK=DK,
num_warps=num_warps, num_stages=num_stages,
)
else:
pO, pM, pL = _get_bufs(NSPLITS, total_q_eff, nh_eff)
done_counter = _get_counter(bs, n_hg)
ckey = (bs, n_hg)
_gen_counter[ckey] += NSPLITS
gen_target = _gen_counter[ckey]
_mla_fused[(bs, NSPLITS, n_hg)](
q_for_kernel, kv_fp8,
pO, pM, pL, output,
qo_for_kernel, kv_indptr,
done_counter,
sms, kv_fp8_scale_val,
q_for_kernel.stride(0), q_for_kernel.stride(1),
kv_fp8.stride(0),
pO.stride(0), pO.stride(1), pO.stride(2),
pM.stride(0), pM.stride(1),
output.stride(0), output.stride(1),
gen_target, nh_eff, qsl_eff, n_hg,
NSPLITS=NSPLITS, BS=BS, HPB=HPB, DK=DK,
num_warps=num_warps, num_stages=num_stages,
)
if qsl > 1:
return output.reshape(bs, qsl, nh, V_DIM).reshape(total_q, nh, V_DIM)
return output
scrolls · 1304 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