submission 716743
ooousay · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1091 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-716743?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:8b5036cf47bd1b109ae93c36278c77ea1252dc0bd3bcf0d225942892b8d02acf
license declaredunknown
license concludedunknown
authorsooousay
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
scores = tl.dot(q_rope, tl.trans(k_rope_bf16))num-warps = 4
num_warps=4, num_stages=1,online-softmax
m_new = tl.maximum(m_i, m_j)split-k
"""bs=4, kv_len=1024 ? Split-K Triton MLA decode with hardcoded strides and indptr elimination.stages = 1
num_warps=4, num_stages=1,Kernel source
submission.py1091 lines
#!POPCORN leaderboard amd-mixed-mla
"""
Auto-generated by build.py ? do not edit directly.
Edit the per-shape kernel.py files and re-run build.py.
"""
# ============================================================
# Constants
# ============================================================
from aiter import dtypes as aiter_dtypes
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 # 576
V_HEAD_DIM = KV_LORA_RANK # 512
SM_SCALE = 1.0 / (QK_HEAD_DIM ** 0.5)
FP8_DTYPE = aiter_dtypes.fp8
# ============================================================
# bs=4, kv_len=1024
# ============================================================
"""bs=4, kv_len=1024 ? Split-K Triton MLA decode with hardcoded strides and indptr elimination.
v33: Hardcode kv_start/q_tok from batch_id, eliminate indptr loads, make strides constexpr.
"""
import torch
import triton
import triton.language as tl
V_CHUNKS = 4 # split 512 V dims into 4 chunks of 128
BLOCK_V_REDUCE = V_HEAD_DIM // V_CHUNKS # 128
@triton.jit
def _mla_stage1_4_1024(
Q_ptr, KV_ptr, KV_scale_ptr,
Partial_ptr, LSE_ptr,
sm_scale: tl.constexpr,
STRIDE_Q_TOK: tl.constexpr, STRIDE_Q_HEAD: tl.constexpr, STRIDE_KV_TOK: tl.constexpr,
BLOCK_KV: tl.constexpr, BLOCK_LORA: tl.constexpr, BLOCK_ROPE: tl.constexpr,
BLOCK_V: tl.constexpr,
HEADS_PER_GROUP: tl.constexpr, NUM_SPLITS: tl.constexpr, NUM_HEADS: tl.constexpr,
BS: tl.constexpr, NUM_ITERS: tl.constexpr,
KV_LEN: tl.constexpr,
):
split_id = tl.program_id(0)
batch_id = tl.program_id(1)
kv_scale = tl.load(KV_scale_ptr)
# Hardcoded: kv_start = batch_id * 1024, no indptr load
kv_start = batch_id * KV_LEN
# Hardcoded: q_tok = batch_id (seqlen=1), no indptr load
q_tok = batch_id
# Load Q split into lora and rope parts
# Bake both sm_scale and kv_scale into Q (loaded once, eliminates per-iteration KV scaling)
h_offs = tl.arange(0, HEADS_PER_GROUP)
lora_offs = tl.arange(0, BLOCK_LORA)
rope_offs = tl.arange(0, BLOCK_ROPE)
q_base = Q_ptr + q_tok * STRIDE_Q_TOK + h_offs[:, None] * STRIDE_Q_HEAD
q_scale = sm_scale * kv_scale
q_lora = (tl.load(q_base + lora_offs[None, :]) * q_scale).to(tl.bfloat16)
q_rope = (tl.load(q_base + (BLOCK_LORA + rope_offs[None, :])) * q_scale).to(tl.bfloat16)
# Accumulators
m_i = tl.full([HEADS_PER_GROUP], float("-inf"), dtype=tl.float32)
l_i = tl.zeros([HEADS_PER_GROUP], dtype=tl.float32)
acc = tl.zeros([HEADS_PER_GROUP, BLOCK_V], dtype=tl.float32)
# tps = KV_LEN // NUM_SPLITS (exact division, no empty splits possible)
kv_base = kv_start + split_id * (KV_LEN // NUM_SPLITS)
v_offs = tl.arange(0, BLOCK_V)
for it in tl.static_range(NUM_ITERS):
tok_offs = tl.arange(0, BLOCK_KV)
tok_idx = kv_base + it * BLOCK_KV + tok_offs
kv_row_ptr = KV_ptr + tok_idx[:, None] * STRIDE_KV_TOK
# Load KV[0:512] as FP8 -> cast to bf16 WITHOUT scaling (scale baked into Q)
kv_shared = tl.load(kv_row_ptr + lora_offs[None, :])
kv_bf16 = kv_shared.to(tl.bfloat16)
# Load K_rope[512:576] as FP8 -> cast to bf16 WITHOUT scaling
k_rope = tl.load(kv_row_ptr + (BLOCK_LORA + rope_offs[None, :]))
k_rope_bf16 = k_rope.to(tl.bfloat16)
# QK^T: kv_scale already baked into Q, so scores are correctly scaled
scores = tl.dot(q_rope, tl.trans(k_rope_bf16))
scores = tl.dot(q_lora, tl.trans(kv_bf16), scores)
# Online softmax
m_j = tl.max(scores, axis=1)
m_new = tl.maximum(m_i, m_j)
alpha = tl.exp(m_i - m_new)
p = tl.exp(scores - m_new[:, None])
l_i = alpha * l_i + tl.sum(p, axis=1)
acc = acc * alpha[:, None]
# PV: V accumulation using unscaled kv_bf16. Scale applied after loop.
acc = tl.dot(p.to(tl.bfloat16), kv_bf16, acc)
m_i = m_new
# Apply kv_scale to V accumulator (once, outside loop)
acc = acc * kv_scale
# Normalize and store partial as bf16 + LSE as f32
norm_acc = acc / tl.maximum(l_i[:, None], 1e-12)
lse = m_i + tl.log(tl.maximum(l_i, 1e-12))
# Partial: [BS, NUM_HEADS, NUM_SPLITS, BLOCK_V] as bf16
partial_base = (Partial_ptr
+ batch_id * (NUM_HEADS * NUM_SPLITS * BLOCK_V)
+ h_offs[:, None] * (NUM_SPLITS * BLOCK_V)
+ split_id * BLOCK_V
+ v_offs[None, :])
tl.store(partial_base, norm_acc.to(tl.bfloat16))
# LSE: [BS, NUM_HEADS, NUM_SPLITS] as f32
lse_base = (LSE_ptr
+ batch_id * (NUM_HEADS * NUM_SPLITS)
+ h_offs * NUM_SPLITS
+ split_id)
tl.store(lse_base, lse)
@triton.jit
def _mla_reduce_vsplit_4_1024(
Partial_ptr, LSE_ptr, O_ptr,
STRIDE_O_TOK: tl.constexpr, STRIDE_O_HEAD: tl.constexpr,
BLOCK_V_FULL: tl.constexpr, BLOCK_V_CHUNK: tl.constexpr,
NUM_SPLITS: tl.constexpr, NUM_HEADS: tl.constexpr,
BS: tl.constexpr,
):
batch_id = tl.program_id(0)
head_id = tl.program_id(1)
v_chunk_id = tl.program_id(2)
v_offs = v_chunk_id * BLOCK_V_CHUNK + tl.arange(0, BLOCK_V_CHUNK)
# Layout: [BS, NUM_HEADS, NUM_SPLITS, ...]
lse_base = LSE_ptr + batch_id * (NUM_HEADS * NUM_SPLITS) + head_id * NUM_SPLITS
partial_base = Partial_ptr + batch_id * (NUM_HEADS * NUM_SPLITS * BLOCK_V_FULL) + head_id * (NUM_SPLITS * BLOCK_V_FULL)
# Pass 1: global max LSE
m_global = tl.full([1], float("-inf"), dtype=tl.float32)
for s in range(NUM_SPLITS):
lse_s = tl.load(lse_base + s)
m_global = tl.maximum(m_global, lse_s)
# Pass 2: weighted sum of normalized partials
acc = tl.zeros([BLOCK_V_CHUNK], dtype=tl.float32)
w_total = tl.zeros([1], dtype=tl.float32)
for s in range(NUM_SPLITS):
lse_s = tl.load(lse_base + s)
w = tl.exp(lse_s - m_global)
partial_s = tl.load(partial_base + s * BLOCK_V_FULL + v_offs).to(tl.float32)
acc += w * partial_s
w_total += w
acc = acc / tl.maximum(w_total, 1e-12)
# Hardcoded: q_tok = batch_id (seqlen=1), no indptr load
tl.store(O_ptr + batch_id * STRIDE_O_TOK + head_id * STRIDE_O_HEAD + v_offs, acc.to(tl.bfloat16))
# === Entry point ===
_state_4_1024 = None
def _run_4_1024(q, kv_data, qo_indptr, kv_indptr, bs, kv_len):
global _state_4_1024
_NUM_SPLITS = 32
_BLOCK_KV = 32
_NUM_HEADS_PER_GROUP = 16
_BS = 4
_KV_LEN = 1024
if _state_4_1024 is None:
_state_4_1024 = {
"partial": torch.empty((bs, NUM_HEADS, _NUM_SPLITS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
"lse": torch.empty((bs, NUM_HEADS, _NUM_SPLITS), dtype=torch.float32, device="cuda"),
"o": torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
}
kv_buf, kv_scale = kv_data["fp8"]
# NUM_ITERS = tokens_per_split / _BLOCK_KV = (1024/32) / 32 = 1
num_iters = _KV_LEN // _NUM_SPLITS // _BLOCK_KV
# Q shape: (4, 16, 576) -> stride_q_tok=16*576=9216, stride_q_head=576
# KV stride: 576
# O shape: (4, 16, 512) -> stride_o_tok=16*512=8192, stride_o_head=512
_mla_stage1_4_1024[(_NUM_SPLITS, bs)](
q, kv_buf, kv_scale,
_state_4_1024["partial"], _state_4_1024["lse"],
SM_SCALE,
STRIDE_Q_TOK=9216, STRIDE_Q_HEAD=576, STRIDE_KV_TOK=576,
BLOCK_KV=_BLOCK_KV, BLOCK_LORA=KV_LORA_RANK, BLOCK_ROPE=QK_ROPE_HEAD_DIM,
BLOCK_V=V_HEAD_DIM,
HEADS_PER_GROUP=_NUM_HEADS_PER_GROUP, NUM_SPLITS=_NUM_SPLITS, NUM_HEADS=NUM_HEADS,
BS=_BS, NUM_ITERS=num_iters,
KV_LEN=_KV_LEN,
num_warps=4, num_stages=1,
)
_mla_reduce_vsplit_4_1024[(bs, NUM_HEADS, V_CHUNKS)](
_state_4_1024["partial"], _state_4_1024["lse"], _state_4_1024["o"],
STRIDE_O_TOK=8192, STRIDE_O_HEAD=512,
BLOCK_V_FULL=V_HEAD_DIM, BLOCK_V_CHUNK=BLOCK_V_REDUCE,
NUM_SPLITS=_NUM_SPLITS, NUM_HEADS=NUM_HEADS,
BS=_BS,
num_warps=4,
)
return _state_4_1024["o"]
# ============================================================
# bs=4, kv_len=8192
# ============================================================
"""bs=4, kv_len=8192 ? Custom Triton MLA decode with split-K + online softmax.
v78: Hardcode shape constants. Eliminate indptr loads, constexpr strides.
Keep 3D grid, num_warps=8, num_stages=2 from v65.
"""
import torch
import triton
import triton.language as tl
@triton.jit
def _mla_stage1_4_8192(
Q_ptr, KV_ptr, KV_scale_ptr,
Partial_ptr, LSE_ptr,
sm_scale: tl.constexpr,
STRIDE_Q_TOK: tl.constexpr, STRIDE_Q_HEAD: tl.constexpr,
STRIDE_KV_TOK: tl.constexpr,
BLOCK_KV: tl.constexpr, BLOCK_LORA: tl.constexpr, BLOCK_ROPE: tl.constexpr,
BLOCK_V: tl.constexpr,
HEADS_PER_GROUP: tl.constexpr, NUM_SPLITS: tl.constexpr, NUM_HEADS: tl.constexpr,
BS: tl.constexpr,
KV_LEN: tl.constexpr, TOKENS_PER_SPLIT: tl.constexpr, NUM_ITERS: tl.constexpr,
):
split_id = tl.program_id(0)
batch_id = tl.program_id(1)
hg = tl.program_id(2)
kv_scale = tl.load(KV_scale_ptr)
# Hardcoded: kv_start = batch_id * 8192, q_tok = batch_id
kv_base = batch_id * KV_LEN + split_id * TOKENS_PER_SPLIT
h_offs = tl.arange(0, HEADS_PER_GROUP)
lora_offs = tl.arange(0, BLOCK_LORA)
rope_offs = tl.arange(0, BLOCK_ROPE)
q_base = Q_ptr + batch_id * STRIDE_Q_TOK + (hg * HEADS_PER_GROUP + h_offs[:, None]) * STRIDE_Q_HEAD
q_scale = sm_scale * kv_scale
q_lora = (tl.load(q_base + lora_offs[None, :]) * q_scale).to(tl.bfloat16)
q_rope = (tl.load(q_base + (BLOCK_LORA + rope_offs[None, :])) * q_scale).to(tl.bfloat16)
m_i = tl.full([HEADS_PER_GROUP], float("-inf"), dtype=tl.float32)
l_i = tl.zeros([HEADS_PER_GROUP], dtype=tl.float32)
acc = tl.zeros([HEADS_PER_GROUP, BLOCK_V], dtype=tl.float32)
v_offs = tl.arange(0, BLOCK_V)
# 8192 / 64 splits / 32 block = 4 exact iterations
for it in range(NUM_ITERS):
tok_offs = tl.arange(0, BLOCK_KV)
tok_idx = kv_base + it * BLOCK_KV + tok_offs
kv_row_ptr = KV_ptr + tok_idx[:, None] * STRIDE_KV_TOK
kv_shared = tl.load(kv_row_ptr + lora_offs[None, :])
kv_bf16 = kv_shared.to(tl.bfloat16)
k_rope = tl.load(kv_row_ptr + (BLOCK_LORA + rope_offs[None, :]))
k_rope_bf16 = k_rope.to(tl.bfloat16)
scores = tl.dot(q_rope, tl.trans(k_rope_bf16))
scores = tl.dot(q_lora, tl.trans(kv_bf16), scores)
m_j = tl.max(scores, axis=1)
m_new = tl.maximum(m_i, m_j)
alpha = tl.exp(m_i - m_new)
p = tl.exp(scores - m_new[:, None])
l_i = alpha * l_i + tl.sum(p, axis=1)
acc = acc * alpha[:, None]
acc = tl.dot(p.to(tl.bfloat16), kv_bf16, acc)
m_i = m_new
acc = acc * kv_scale
norm_acc = acc / tl.maximum(l_i[:, None], 1e-12)
lse = m_i + tl.log(tl.maximum(l_i, 1e-12))
partial_base = (Partial_ptr
+ batch_id * (NUM_HEADS * NUM_SPLITS * BLOCK_V)
+ (hg * HEADS_PER_GROUP + h_offs[:, None]) * (NUM_SPLITS * BLOCK_V)
+ split_id * BLOCK_V
+ v_offs[None, :])
tl.store(partial_base, norm_acc.to(tl.bfloat16))
lse_base = (LSE_ptr
+ batch_id * (NUM_HEADS * NUM_SPLITS)
+ (hg * HEADS_PER_GROUP + h_offs) * NUM_SPLITS
+ split_id)
tl.store(lse_base, lse)
@triton.jit
def _mla_reduce_4_8192(
Partial_ptr, LSE_ptr, O_ptr,
STRIDE_O_TOK: tl.constexpr, STRIDE_O_HEAD: tl.constexpr,
BLOCK_V: tl.constexpr, NUM_SPLITS: tl.constexpr, NUM_HEADS: tl.constexpr,
BS: tl.constexpr,
):
batch_id = tl.program_id(0)
head_id = tl.program_id(1)
v_offs = tl.arange(0, BLOCK_V)
lse_base = LSE_ptr + batch_id * (NUM_HEADS * NUM_SPLITS) + head_id * NUM_SPLITS
partial_base = Partial_ptr + batch_id * (NUM_HEADS * NUM_SPLITS * BLOCK_V) + head_id * (NUM_SPLITS * BLOCK_V)
m_global = tl.full([1], float("-inf"), dtype=tl.float32)
for s in range(NUM_SPLITS):
lse_s = tl.load(lse_base + s)
m_global = tl.maximum(m_global, lse_s)
acc = tl.zeros([BLOCK_V], dtype=tl.float32)
w_total = tl.zeros([1], dtype=tl.float32)
for s in range(NUM_SPLITS):
lse_s = tl.load(lse_base + s)
w = tl.exp(lse_s - m_global)
partial_s = tl.load(partial_base + s * BLOCK_V + v_offs)
acc += w * partial_s
w_total += w
acc = acc / tl.maximum(w_total, 1e-12)
tl.store(O_ptr + batch_id * STRIDE_O_TOK + head_id * STRIDE_O_HEAD + v_offs, acc.to(tl.bfloat16))
_state_4_8192 = None
def _run_4_8192(q, kv_data, qo_indptr, kv_indptr, bs, kv_len):
global _state_4_8192
_NUM_SPLITS = 64
_BLOCK_KV = 32
_NUM_HEADS_PER_GROUP = 16
_BS = 4
_NUM_ITERS = 4 # 8192 / 64 / 32
if _state_4_8192 is None:
_state_4_8192 = {
"partial": torch.empty((bs, NUM_HEADS, _NUM_SPLITS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
"lse": torch.empty((bs, NUM_HEADS, _NUM_SPLITS), dtype=torch.float32, device="cuda"),
"o": torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
}
kv_buf, kv_scale = kv_data["fp8"]
nhg = NUM_HEADS // _NUM_HEADS_PER_GROUP # = 1
_mla_stage1_4_8192[(_NUM_SPLITS, bs, nhg)](
q, kv_buf, kv_scale,
_state_4_8192["partial"], _state_4_8192["lse"],
SM_SCALE,
STRIDE_Q_TOK=9216, STRIDE_Q_HEAD=576,
STRIDE_KV_TOK=576,
BLOCK_KV=_BLOCK_KV, BLOCK_LORA=KV_LORA_RANK, BLOCK_ROPE=QK_ROPE_HEAD_DIM,
BLOCK_V=V_HEAD_DIM,
HEADS_PER_GROUP=_NUM_HEADS_PER_GROUP, NUM_SPLITS=_NUM_SPLITS, NUM_HEADS=NUM_HEADS,
BS=_BS,
KV_LEN=8192, TOKENS_PER_SPLIT=128, NUM_ITERS=_NUM_ITERS,
num_warps=8, num_stages=2,
)
_mla_reduce_4_8192[(bs, NUM_HEADS)](
_state_4_8192["partial"], _state_4_8192["lse"], _state_4_8192["o"],
STRIDE_O_TOK=8192, STRIDE_O_HEAD=512,
BLOCK_V=V_HEAD_DIM, NUM_SPLITS=_NUM_SPLITS, NUM_HEADS=NUM_HEADS,
BS=_BS,
num_warps=4,
)
return _state_4_8192["o"]
# ============================================================
# bs=32, kv_len=1024
# ============================================================
"""bs=32, kv_len=1024 ? Split-K Triton MLA decode with XCD remapping.
v12: Fully hardcoded for (bs=32, kv=1024, seqlen=1). No indptr loads, constexpr
strides, no empty-split check. Keeps XCD remap + static_range from v8.
"""
import torch
import triton
import triton.language as tl
@triton.jit
def remap_xcd_32_1024(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):
pids_per_xcd = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS
tall_xcds = GRID_MN % NUM_XCDS
tall_xcds = NUM_XCDS if tall_xcds == 0 else tall_xcds
xcd = pid % NUM_XCDS
local_pid = pid // NUM_XCDS
if xcd < tall_xcds:
pid = xcd * pids_per_xcd + local_pid
else:
pid = (
tall_xcds * pids_per_xcd
+ (xcd - tall_xcds) * (pids_per_xcd - 1)
+ local_pid
)
return pid
@triton.jit
def _mla_stage1_32_1024(
Q_ptr, KV_ptr, KV_scale_ptr,
Partial_ptr, LSE_ptr,
sm_scale: tl.constexpr,
STRIDE_Q_TOK: tl.constexpr, STRIDE_Q_HEAD: tl.constexpr,
STRIDE_KV_TOK: tl.constexpr,
BLOCK_KV: tl.constexpr, BLOCK_LORA: tl.constexpr, BLOCK_ROPE: tl.constexpr,
BLOCK_V: tl.constexpr,
HEADS_PER_GROUP: tl.constexpr, NUM_SPLITS: tl.constexpr, NUM_HEADS: tl.constexpr,
NUM_ITERS: tl.constexpr,
KV_LEN: tl.constexpr, TOKENS_PER_SPLIT: tl.constexpr,
GRID_TOTAL: tl.constexpr,
):
pid = tl.program_id(0)
pid = remap_xcd_32_1024(pid, GRID_TOTAL)
split_id = pid % NUM_SPLITS
batch_id = pid // NUM_SPLITS
kv_scale = tl.load(KV_scale_ptr)
# Hardcoded: kv_start = batch_id * 1024, q_tok = batch_id (seqlen=1)
kv_base = batch_id * KV_LEN + split_id * TOKENS_PER_SPLIT
h_offs = tl.arange(0, HEADS_PER_GROUP)
lora_offs = tl.arange(0, BLOCK_LORA)
rope_offs = tl.arange(0, BLOCK_ROPE)
q_base = Q_ptr + batch_id * STRIDE_Q_TOK + h_offs[:, None] * STRIDE_Q_HEAD
q_scale = sm_scale * kv_scale
q_lora = (tl.load(q_base + lora_offs[None, :]) * q_scale).to(tl.bfloat16)
q_rope = (tl.load(q_base + (BLOCK_LORA + rope_offs[None, :])) * q_scale).to(tl.bfloat16)
m_i = tl.full([HEADS_PER_GROUP], float("-inf"), dtype=tl.float32)
l_i = tl.zeros([HEADS_PER_GROUP], dtype=tl.float32)
acc = tl.zeros([HEADS_PER_GROUP, BLOCK_V], dtype=tl.float32)
v_offs = tl.arange(0, BLOCK_V)
# 1024 / 8 splits / 32 block = 4 exact iterations
for it in tl.static_range(NUM_ITERS):
tok_offs = tl.arange(0, BLOCK_KV)
tok_idx = kv_base + it * BLOCK_KV + tok_offs
kv_row_ptr = KV_ptr + tok_idx[:, None] * STRIDE_KV_TOK
kv_shared = tl.load(kv_row_ptr + lora_offs[None, :])
kv_bf16 = kv_shared.to(tl.bfloat16)
k_rope = tl.load(kv_row_ptr + (BLOCK_LORA + rope_offs[None, :]))
k_rope_bf16 = k_rope.to(tl.bfloat16)
scores = tl.dot(q_rope, tl.trans(k_rope_bf16))
scores = tl.dot(q_lora, tl.trans(kv_bf16), scores)
m_j = tl.max(scores, axis=1)
m_new = tl.maximum(m_i, m_j)
alpha = tl.exp(m_i - m_new)
p = tl.exp(scores - m_new[:, None])
l_i = alpha * l_i + tl.sum(p, axis=1)
acc = acc * alpha[:, None]
acc = tl.dot(p.to(tl.bfloat16), kv_bf16, acc)
m_i = m_new
acc = acc * kv_scale
norm_acc = acc / tl.maximum(l_i[:, None], 1e-12)
lse = m_i + tl.log(tl.maximum(l_i, 1e-12))
partial_base = (Partial_ptr
+ batch_id * (NUM_HEADS * NUM_SPLITS * BLOCK_V)
+ h_offs[:, None] * (NUM_SPLITS * BLOCK_V)
+ split_id * BLOCK_V
+ v_offs[None, :])
tl.store(partial_base, norm_acc.to(tl.bfloat16))
lse_base = (LSE_ptr
+ batch_id * (NUM_HEADS * NUM_SPLITS)
+ h_offs * NUM_SPLITS
+ split_id)
tl.store(lse_base, lse)
@triton.jit
def _mla_reduce_32_1024(
Partial_ptr, LSE_ptr, O_ptr,
STRIDE_O_TOK: tl.constexpr, STRIDE_O_HEAD: tl.constexpr,
BLOCK_V: tl.constexpr, NUM_SPLITS: tl.constexpr, NUM_HEADS: tl.constexpr,
GRID_TOTAL: tl.constexpr,
):
pid = tl.program_id(0)
pid = remap_xcd_32_1024(pid, GRID_TOTAL)
batch_id = pid // NUM_HEADS
head_id = pid % NUM_HEADS
v_offs = tl.arange(0, BLOCK_V)
lse_base = LSE_ptr + batch_id * (NUM_HEADS * NUM_SPLITS) + head_id * NUM_SPLITS
partial_base = Partial_ptr + batch_id * (NUM_HEADS * NUM_SPLITS * BLOCK_V) + head_id * (NUM_SPLITS * BLOCK_V)
m_global = tl.full([1], float("-inf"), dtype=tl.float32)
for s in range(NUM_SPLITS):
lse_s = tl.load(lse_base + s)
m_global = tl.maximum(m_global, lse_s)
acc = tl.zeros([BLOCK_V], dtype=tl.float32)
w_total = tl.zeros([1], dtype=tl.float32)
for s in range(NUM_SPLITS):
lse_s = tl.load(lse_base + s)
w = tl.exp(lse_s - m_global)
partial_s = tl.load(partial_base + s * BLOCK_V + v_offs).to(tl.float32)
acc += w * partial_s
w_total += w
acc = acc / tl.maximum(w_total, 1e-12)
# Hardcoded: q_tok = batch_id (seqlen=1)
tl.store(O_ptr + batch_id * STRIDE_O_TOK + head_id * STRIDE_O_HEAD + v_offs, acc.to(tl.bfloat16))
_state_32_1024 = None
def _run_32_1024(q, kv_data, qo_indptr, kv_indptr, bs, kv_len):
global _state_32_1024
_NUM_SPLITS = 8
_BLOCK_KV = 32
_NUM_HEADS_PER_GROUP = 16
_NUM_ITERS = 4 # 1024 / 8 / 32
if _state_32_1024 is None:
_state_32_1024 = {
"partial": torch.empty((bs, NUM_HEADS, _NUM_SPLITS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
"lse": torch.empty((bs, NUM_HEADS, _NUM_SPLITS), dtype=torch.float32, device="cuda"),
"o": torch.empty((bs, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
}
kv_buf, kv_scale = kv_data["fp8"]
stage1_grid = _NUM_SPLITS * bs # 256
reduce_grid = bs * NUM_HEADS # 512
_mla_stage1_32_1024[(stage1_grid,)](
q, kv_buf, kv_scale,
_state_32_1024["partial"], _state_32_1024["lse"],
SM_SCALE,
STRIDE_Q_TOK=9216, STRIDE_Q_HEAD=576,
STRIDE_KV_TOK=576,
BLOCK_KV=_BLOCK_KV, BLOCK_LORA=KV_LORA_RANK, BLOCK_ROPE=QK_ROPE_HEAD_DIM,
BLOCK_V=V_HEAD_DIM,
HEADS_PER_GROUP=_NUM_HEADS_PER_GROUP, NUM_SPLITS=_NUM_SPLITS, NUM_HEADS=NUM_HEADS,
NUM_ITERS=_NUM_ITERS,
KV_LEN=1024, TOKENS_PER_SPLIT=128,
GRID_TOTAL=stage1_grid,
num_warps=4, num_stages=1,
)
_mla_reduce_32_1024[(reduce_grid,)](
_state_32_1024["partial"], _state_32_1024["lse"], _state_32_1024["o"],
STRIDE_O_TOK=8192, STRIDE_O_HEAD=512,
BLOCK_V=V_HEAD_DIM, NUM_SPLITS=_NUM_SPLITS, NUM_HEADS=NUM_HEADS,
GRID_TOTAL=reduce_grid,
num_warps=4,
)
return _state_32_1024["o"]
# ============================================================
# bs=32, kv_len=8192
# ============================================================
"""aiter a16w8 with page_size=8 ? scheduler sees 1/8 sequence length."""
import torch
import aiter
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
_state_32_8192 = None
def _run_32_8192(q, kv_data, qo_indptr, kv_indptr, bs, kv_len):
global _state_32_8192
PAGE_SIZE = 8
NUM_KV_SPLITS = 32
nq, nkv, q_seq_len, dv = NUM_HEADS, NUM_KV_HEADS, 1, V_HEAD_DIM
total_q = bs * q_seq_len
if _state_32_8192 is None:
pages_per_batch = kv_len // PAGE_SIZE
total_pages = bs * pages_per_batch
kv_indptr_pages = torch.zeros(bs + 1, dtype=torch.int32, device="cuda")
for i in range(bs):
kv_indptr_pages[i + 1] = kv_indptr_pages[i] + pages_per_batch
kv_indices = torch.arange(total_pages, dtype=torch.int32, device="cuda")
kv_last_page_len = torch.full((bs,), PAGE_SIZE, dtype=torch.int32, device="cuda")
q_dtype, kv_dtype = torch.bfloat16, FP8_DTYPE
info = get_mla_metadata_info_v1(
bs, q_seq_len, nq, q_dtype, kv_dtype,
is_sparse=False, fast_mode=False,
num_kv_splits=NUM_KV_SPLITS, intra_batch_mode=False,
)
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_pages, kv_last_page_len,
nq // nkv, nkv, 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=q_seq_len, uni_seqlen_qo=q_seq_len,
fast_mode=False, max_split_per_batch=NUM_KV_SPLITS,
intra_batch_mode=False, dtype_q=q_dtype, dtype_kv=kv_dtype,
)
num_partial = reduce_partial_map.size(0)
_state_32_8192 = {
"kv_indices": kv_indices,
"kv_last_page_len": kv_last_page_len,
"kv_indptr_pages": kv_indptr_pages,
"work_metadata": 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,
"logits": torch.empty((num_partial * q_seq_len, 1, nq, dv), dtype=torch.float32, device="cuda"),
"attn_lse": torch.empty((num_partial * q_seq_len, 1, nq, 1), dtype=torch.float32, device="cuda"),
"o": torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda"),
}
kv_buf, kv_scale = kv_data["fp8"]
kv_4d = kv_buf.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_buf.shape[-1])
aiter.mla_decode_stage1_asm_fwd(
q.view(-1, NUM_HEADS, QK_HEAD_DIM),
kv_4d,
qo_indptr, _state_32_8192["kv_indptr_pages"],
_state_32_8192["kv_indices"], _state_32_8192["kv_last_page_len"],
None,
_state_32_8192["work_metadata"], _state_32_8192["work_indptr"], _state_32_8192["work_info_set"],
1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
_state_32_8192["logits"], _state_32_8192["attn_lse"], _state_32_8192["o"],
None, kv_scale,
)
aiter.mla_reduce_v1(
_state_32_8192["logits"], _state_32_8192["attn_lse"],
_state_32_8192["reduce_indptr"], _state_32_8192["reduce_final_map"], _state_32_8192["reduce_partial_map"],
1, _state_32_8192["o"], None,
)
return _state_32_8192["o"]
# ============================================================
# bs=64, kv_len=1024
# ============================================================
"""bs=64, kv_len=1024 ? aiter a16w8 with page_size=2.
v22: Switch from custom Triton to aiter ASM with page_size=2.
"""
import torch
import aiter
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
_state_64_1024 = None
def _run_64_1024(q, kv_data, qo_indptr, kv_indptr, bs, kv_len):
global _state_64_1024
PAGE_SIZE = 2
NUM_KV_SPLITS = 32
nq, nkv, q_seq_len, dv = NUM_HEADS, NUM_KV_HEADS, 1, V_HEAD_DIM
total_q = bs * q_seq_len
if _state_64_1024 is None:
pages_per_batch = kv_len // PAGE_SIZE
total_pages = bs * pages_per_batch
kv_indptr_pages = torch.zeros(bs + 1, dtype=torch.int32, device="cuda")
for i in range(bs):
kv_indptr_pages[i + 1] = kv_indptr_pages[i] + pages_per_batch
kv_indices = torch.arange(total_pages, dtype=torch.int32, device="cuda")
kv_last_page_len = torch.full((bs,), PAGE_SIZE, dtype=torch.int32, device="cuda")
q_dtype, kv_dtype = torch.bfloat16, FP8_DTYPE
info = get_mla_metadata_info_v1(
bs, q_seq_len, nq, q_dtype, kv_dtype,
is_sparse=False, fast_mode=False,
num_kv_splits=NUM_KV_SPLITS, intra_batch_mode=False,
)
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_pages, kv_last_page_len,
nq // nkv, nkv, 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=q_seq_len, uni_seqlen_qo=q_seq_len,
fast_mode=False, max_split_per_batch=NUM_KV_SPLITS,
intra_batch_mode=False, dtype_q=q_dtype, dtype_kv=kv_dtype,
)
num_partial = reduce_partial_map.size(0)
_state_64_1024 = {
"kv_indices": kv_indices,
"kv_last_page_len": kv_last_page_len,
"kv_indptr_pages": kv_indptr_pages,
"work_metadata": 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,
"logits": torch.empty((num_partial * q_seq_len, 1, nq, dv), dtype=torch.float32, device="cuda"),
"attn_lse": torch.empty((num_partial * q_seq_len, 1, nq, 1), dtype=torch.float32, device="cuda"),
"o": torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda"),
}
kv_buf, kv_scale = kv_data["fp8"]
kv_4d = kv_buf.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_buf.shape[-1])
aiter.mla_decode_stage1_asm_fwd(
q.view(-1, NUM_HEADS, QK_HEAD_DIM),
kv_4d,
qo_indptr, _state_64_1024["kv_indptr_pages"],
_state_64_1024["kv_indices"], _state_64_1024["kv_last_page_len"],
None,
_state_64_1024["work_metadata"], _state_64_1024["work_indptr"], _state_64_1024["work_info_set"],
1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
_state_64_1024["logits"], _state_64_1024["attn_lse"], _state_64_1024["o"],
None, kv_scale,
)
aiter.mla_reduce_v1(
_state_64_1024["logits"], _state_64_1024["attn_lse"],
_state_64_1024["reduce_indptr"], _state_64_1024["reduce_final_map"], _state_64_1024["reduce_partial_map"],
1, _state_64_1024["o"], None,
)
return _state_64_1024["o"]
# ============================================================
# bs=64, kv_len=8192
# ============================================================
"""aiter a16w8 page_size=8 intra_batch_mode=True.
Hypothesis: with bs=64 all sequences same kv_len=8192, intra_batch_mode
gives scheduler uniform work distribution ? better CU utilization.
"""
import torch
import aiter
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
_state_64_8192 = None
def _run_64_8192(q, kv_data, qo_indptr, kv_indptr, bs, kv_len):
global _state_64_8192
PAGE_SIZE = 8
NUM_KV_SPLITS = 32
nq, nkv, q_seq_len, dv = NUM_HEADS, NUM_KV_HEADS, 1, V_HEAD_DIM
total_q = bs * q_seq_len
if _state_64_8192 is None:
pages_per_batch = kv_len // PAGE_SIZE
total_pages = bs * pages_per_batch
kv_indptr_pages = torch.zeros(bs + 1, dtype=torch.int32, device="cuda")
for i in range(bs):
kv_indptr_pages[i + 1] = kv_indptr_pages[i] + pages_per_batch
kv_indices = torch.arange(total_pages, dtype=torch.int32, device="cuda")
kv_last_page_len = torch.full((bs,), PAGE_SIZE, dtype=torch.int32, device="cuda")
q_dtype, kv_dtype = torch.bfloat16, FP8_DTYPE
info = get_mla_metadata_info_v1(
bs, q_seq_len, 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="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_pages, kv_last_page_len,
nq // nkv, nkv, 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=q_seq_len, uni_seqlen_qo=q_seq_len,
fast_mode=False, max_split_per_batch=NUM_KV_SPLITS,
intra_batch_mode=True, dtype_q=q_dtype, dtype_kv=kv_dtype,
)
num_partial = reduce_partial_map.size(0)
_state_64_8192 = {
"kv_indices": kv_indices,
"kv_last_page_len": kv_last_page_len,
"kv_indptr_pages": kv_indptr_pages,
"work_metadata": 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,
"logits": torch.empty((num_partial * q_seq_len, 1, nq, dv), dtype=torch.float32, device="cuda"),
"attn_lse": torch.empty((num_partial * q_seq_len, 1, nq, 1), dtype=torch.float32, device="cuda"),
"o": torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda"),
}
kv_buf, kv_scale = kv_data["fp8"]
kv_4d = kv_buf.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_buf.shape[-1])
aiter.mla_decode_stage1_asm_fwd(
q.view(-1, NUM_HEADS, QK_HEAD_DIM),
kv_4d,
qo_indptr, _state_64_8192["kv_indptr_pages"],
_state_64_8192["kv_indices"], _state_64_8192["kv_last_page_len"],
None,
_state_64_8192["work_metadata"], _state_64_8192["work_indptr"], _state_64_8192["work_info_set"],
1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
_state_64_8192["logits"], _state_64_8192["attn_lse"], _state_64_8192["o"],
None, kv_scale,
)
aiter.mla_reduce_v1(
_state_64_8192["logits"], _state_64_8192["attn_lse"],
_state_64_8192["reduce_indptr"], _state_64_8192["reduce_final_map"], _state_64_8192["reduce_partial_map"],
1, _state_64_8192["o"], None,
)
return _state_64_8192["o"]
# ============================================================
# bs=256, kv_len=1024
# ============================================================
"""bs=256, kv_len=1024 ? aiter a16w8 with page_size=2.
v6: Fix page_size metadata ? kv_indptr in pages, page-based indices.
"""
import torch
import aiter
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
_state_256_1024 = None
def _run_256_1024(q, kv_data, qo_indptr, kv_indptr, bs, kv_len):
global _state_256_1024
PAGE_SIZE = 2
NUM_KV_SPLITS = 32
MODE = "a16w8"
nq, nkv, q_seq_len, dv = NUM_HEADS, NUM_KV_HEADS, 1, V_HEAD_DIM
total_q = bs * q_seq_len
if _state_256_1024 is None:
pages_per_batch = kv_len // PAGE_SIZE # 1024 / 2 = 512
total_pages = bs * pages_per_batch
kv_indptr_pages = torch.zeros(bs + 1, dtype=torch.int32, device="cuda")
for i in range(bs):
kv_indptr_pages[i + 1] = kv_indptr_pages[i] + pages_per_batch
kv_indices = torch.arange(total_pages, dtype=torch.int32, device="cuda")
kv_last_page_len = torch.full((bs,), PAGE_SIZE, dtype=torch.int32, device="cuda")
q_dtype, kv_dtype = torch.bfloat16, FP8_DTYPE
info = get_mla_metadata_info_v1(
bs, q_seq_len, nq, q_dtype, kv_dtype,
is_sparse=False, fast_mode=True,
num_kv_splits=NUM_KV_SPLITS, intra_batch_mode=False,
)
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_pages, kv_last_page_len,
nq // nkv, nkv, 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=q_seq_len, uni_seqlen_qo=q_seq_len,
fast_mode=True, max_split_per_batch=NUM_KV_SPLITS,
intra_batch_mode=False, dtype_q=q_dtype, dtype_kv=kv_dtype,
)
num_partial = reduce_partial_map.size(0)
_state_256_1024 = {
"kv_indices": kv_indices,
"kv_last_page_len": kv_last_page_len,
"kv_indptr_pages": kv_indptr_pages,
"work_metadata": 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,
"logits": torch.empty((num_partial * q_seq_len, 1, nq, dv), dtype=torch.float32, device="cuda"),
"attn_lse": torch.empty((num_partial * q_seq_len, 1, nq, 1), dtype=torch.float32, device="cuda"),
"o": torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda"),
}
kv_buf, kv_scale = kv_data["fp8"]
kv_4d = kv_buf.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_buf.shape[-1])
aiter.mla_decode_stage1_asm_fwd(
q.view(-1, NUM_HEADS, QK_HEAD_DIM),
kv_4d,
qo_indptr, _state_256_1024["kv_indptr_pages"],
_state_256_1024["kv_indices"], _state_256_1024["kv_last_page_len"],
None,
_state_256_1024["work_metadata"], _state_256_1024["work_indptr"], _state_256_1024["work_info_set"],
1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
_state_256_1024["logits"], _state_256_1024["attn_lse"], _state_256_1024["o"],
None, kv_scale,
)
aiter.mla_reduce_v1(
_state_256_1024["logits"], _state_256_1024["attn_lse"],
_state_256_1024["reduce_indptr"], _state_256_1024["reduce_final_map"], _state_256_1024["reduce_partial_map"],
1, _state_256_1024["o"], None,
)
return _state_256_1024["o"]
# ============================================================
# bs=256, kv_len=8192
# ============================================================
"""aiter a16w8 page_size=8 splits=64 intra_batch_mode=True.
Hypothesis: with bs=256 all sequences same length, intra_batch_mode gives
scheduler uniform work distribution ? better CU utilization.
"""
import torch
import aiter
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
_state_256_8192 = None
def _run_256_8192(q, kv_data, qo_indptr, kv_indptr, bs, kv_len):
global _state_256_8192
PAGE_SIZE = 8
NUM_KV_SPLITS = 64
nq, nkv, q_seq_len, dv = NUM_HEADS, NUM_KV_HEADS, 1, V_HEAD_DIM
total_q = bs * q_seq_len
if _state_256_8192 is None:
pages_per_batch = kv_len // PAGE_SIZE
total_pages = bs * pages_per_batch
kv_indptr_pages = torch.zeros(bs + 1, dtype=torch.int32, device="cuda")
for i in range(bs):
kv_indptr_pages[i + 1] = kv_indptr_pages[i] + pages_per_batch
kv_indices = torch.arange(total_pages, dtype=torch.int32, device="cuda")
kv_last_page_len = torch.full((bs,), PAGE_SIZE, dtype=torch.int32, device="cuda")
q_dtype, kv_dtype = torch.bfloat16, FP8_DTYPE
info = get_mla_metadata_info_v1(
bs, q_seq_len, nq, q_dtype, kv_dtype,
is_sparse=False, fast_mode=True,
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_pages, kv_last_page_len,
nq // nkv, nkv, 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=q_seq_len, uni_seqlen_qo=q_seq_len,
fast_mode=True, max_split_per_batch=NUM_KV_SPLITS,
intra_batch_mode=True, dtype_q=q_dtype, dtype_kv=kv_dtype,
)
num_partial = reduce_partial_map.size(0)
_state_256_8192 = {
"kv_indices": kv_indices,
"kv_last_page_len": kv_last_page_len,
"kv_indptr_pages": kv_indptr_pages,
"work_metadata": 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,
"logits": torch.empty((num_partial * q_seq_len, 1, nq, dv), dtype=torch.float32, device="cuda"),
"attn_lse": torch.empty((num_partial * q_seq_len, 1, nq, 1), dtype=torch.float32, device="cuda"),
"o": torch.empty((total_q, nq, dv), dtype=torch.bfloat16, device="cuda"),
}
kv_buf, kv_scale = kv_data["fp8"]
kv_4d = kv_buf.view(-1, PAGE_SIZE, NUM_KV_HEADS, kv_buf.shape[-1])
aiter.mla_decode_stage1_asm_fwd(
q.view(-1, NUM_HEADS, QK_HEAD_DIM),
kv_4d,
qo_indptr, _state_256_8192["kv_indptr_pages"],
_state_256_8192["kv_indices"], _state_256_8192["kv_last_page_len"],
None,
_state_256_8192["work_metadata"], _state_256_8192["work_indptr"], _state_256_8192["work_info_set"],
1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
_state_256_8192["logits"], _state_256_8192["attn_lse"], _state_256_8192["o"],
None, kv_scale,
)
aiter.mla_reduce_v1(
_state_256_8192["logits"], _state_256_8192["attn_lse"],
_state_256_8192["reduce_indptr"], _state_256_8192["reduce_final_map"], _state_256_8192["reduce_partial_map"],
1, _state_256_8192["o"], None,
)
return _state_256_8192["o"]
# ============================================================
# Dispatch
# ============================================================
from task import input_t, output_t
_DISPATCH = {
(4, 1024): _run_4_1024,
(4, 8192): _run_4_8192,
(32, 1024): _run_32_1024,
(32, 8192): _run_32_8192,
(64, 1024): _run_64_1024,
(64, 8192): _run_64_8192,
(256, 1024): _run_256_1024,
(256, 8192): _run_256_8192,
}
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
bs = config["batch_size"]
kv_len = config["kv_seq_len"]
return _DISPATCH[(bs, kv_len)](q, kv_data, qo_indptr, kv_indptr, bs, kv_len)
scrolls · 1091 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