submission 703488
romepen788 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 652 lines, June 9 Researcher Reciprocity License v1.0.
mla_test.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-703488?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:c20062c97643c78f7f21d3dfd0afda0f5a28ec9aaf8e4cbbb6de93c370ecd908
license declaredunknown
license concludedunknown
authorsromepen788
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
scores += tl.dot(q_tile.to(k_tile.dtype), tl.trans(k_tile))num-warps = 8
num_warps=8, num_stages=3,online-softmax
m_new = tl.maximum(m_i, row_max)persistent-kernel
- S4, S6, S7, S8: aiter persistent fp8 (v143)split-k
def _reduce_splitk_parallel_s1(stages = 3
num_warps=8, num_stages=3,Kernel source
mla_test.py652 lines
"""
mla_test: merged best kernels from mla_experiments_1235
- S1 (batch=4, kv=1024): triton s1_v45 parallel reduce w1
- S2 (batch=4, kv=8192): triton s2_v57 parallel reduce w2
- S3 (batch=32, kv=1024): triton s3_v37 serial reduce w1
- S5 (batch=64, kv=1024): triton s5_v15 no-vtile rv512
- S4, S6, S7, S8: aiter persistent fp8 (v143)
"""
import torch
import triton
import triton.language as tl
import aiter
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes
from aiter import get_mla_metadata_info_v1, get_mla_metadata_v1
from aiter.mla import get_meta_param, _fwd_kernel_stage2_asm
NUM_HEADS: tl.constexpr = 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)
PAGE_SIZE = 1
FP8_DTYPE = aiter_dtypes.fp8
_cache = {}
# ============================================================
# S1 triton kernels (batch=4, kv=1024)
# splits=8, SPLIT_LEN=128, V_BLOCK=128, parallel reduce w1
# ============================================================
S1_NUM_SPLITS = 8
S1_SPLIT_LEN = 128
S1_KV_SEQ_LEN = 1024
S1_V_BLOCK = 128
@triton.jit
def _flash_decode_fp8_s1_exact_vtile(
Q_ptr, KV_ptr, Mid_O, Mid_lse,
stride_kv: tl.int64, kv_scale,
sm_scale: tl.constexpr, QK_DIM: tl.constexpr, V_DIM: tl.constexpr,
BLOCK_K: tl.constexpr, NH: tl.constexpr, V_BLOCK_C: tl.constexpr,
NUM_SPLITS_C: tl.constexpr, SPLIT_LEN_C: tl.constexpr, KV_SEQ_LEN_C: tl.constexpr,
):
batch_id = tl.program_id(0)
split_id = tl.program_id(1)
v_block_id = tl.program_id(2)
out_idx = batch_id * NUM_SPLITS_C + split_id
kv_pos = batch_id * KV_SEQ_LEN_C + split_id * SPLIT_LEN_C
tl.multiple_of(kv_pos, 128)
h_offs = tl.arange(0, NH)
n_offs = tl.arange(0, SPLIT_LEN_C)
v_start = v_block_id * V_BLOCK_C
v_offs = v_start + tl.arange(0, V_BLOCK_C)
v_mask = v_offs < V_DIM
q_base = batch_id * NH * QK_DIM
tl.multiple_of(q_base, 128)
tl.multiple_of(stride_kv, 32)
tl.assume(stride_kv > 0)
acc = tl.zeros([NH, V_BLOCK_C], dtype=tl.float32)
m_i = tl.full([NH], float("-inf"), dtype=tl.float32)
l_i = tl.zeros([NH], dtype=tl.float32)
score_scale = sm_scale * kv_scale
scores = tl.zeros([NH, SPLIT_LEN_C], dtype=tl.float32)
for k_start in range(0, QK_DIM, BLOCK_K):
k_offs = tl.arange(0, BLOCK_K)
k_mask = (k_start + k_offs) < QK_DIM
q_tile = tl.load(Q_ptr + q_base + h_offs[:, None] * QK_DIM + k_start + k_offs[None, :], mask=k_mask[None, :], other=0.0)
k_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + k_start + k_offs[None, :], mask=k_mask[None, :], other=0.0)
scores += tl.dot(q_tile.to(k_tile.dtype), tl.trans(k_tile))
scores *= score_scale
row_max = tl.max(scores, axis=1)
m_new = tl.maximum(m_i, row_max)
alpha = tl.exp(m_i - m_new)
l_i = l_i * alpha
exp_scores = tl.exp(scores - m_new[:, None])
l_i += tl.sum(exp_scores, axis=1)
acc = acc * alpha[:, None]
v_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + v_offs[None, :], mask=v_mask[None, :], other=0.0)
acc += tl.dot(exp_scores.to(v_tile.dtype), v_tile)
m_i = m_new
acc = (acc * kv_scale) / l_i[:, None]
lse_vals = m_i + tl.log(l_i)
tl.store(Mid_O + out_idx * NH * V_DIM + h_offs[:, None] * V_DIM + v_offs[None, :], acc.to(tl.bfloat16), mask=v_mask[None, :])
tl.store(Mid_lse + out_idx * NH + h_offs, lse_vals)
@triton.jit
def _reduce_splitk_parallel_s1(
Mid_O, Mid_lse, O_ptr,
NUM_SPLITS_C: tl.constexpr, V_DIM: tl.constexpr, NH: tl.constexpr, BLOCK_V: tl.constexpr,
):
batch_id = tl.program_id(0)
head_id = tl.program_id(1)
v_block = tl.program_id(2)
v_offs = v_block * BLOCK_V + tl.arange(0, BLOCK_V)
v_mask = v_offs < V_DIM
s_offs = tl.arange(0, NUM_SPLITS_C)
all_lse = tl.load(Mid_lse + (batch_id * NUM_SPLITS_C + s_offs) * NH + head_id)
m_global = tl.max(all_lse, axis=0)
alphas = tl.exp(all_lse - m_global)
l_total = tl.sum(alphas, axis=0)
weights = alphas / l_total
all_partials = tl.load(
Mid_O + (batch_id * NUM_SPLITS_C + s_offs[:, None]) * NH * V_DIM + head_id * V_DIM + v_offs[None, :],
mask=v_mask[None, :], other=0.0,
).to(tl.float32)
result = tl.sum(weights[:, None] * all_partials, axis=0)
out_base = batch_id * NH * V_DIM + head_id * V_DIM
tl.store(O_ptr + out_base + v_offs, result.to(tl.bfloat16), mask=v_mask)
# ============================================================
# S2 triton kernels (batch=4, kv=8192)
# splits=16, SPLIT_LEN=512, V_BLOCK=128, parallel reduce w2
# ============================================================
S2_NUM_SPLITS = 16
S2_SPLIT_LEN = 512
S2_KV_SEQ_LEN = 8192
S2_V_BLOCK = 128
@triton.jit
def _flash_decode_fp8_s2_exact_vtile(
Q_ptr, KV_ptr, Mid_O, Mid_lse,
stride_kv: tl.int64, kv_scale,
sm_scale: tl.constexpr, QK_DIM: tl.constexpr, V_DIM: tl.constexpr,
BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, NH: tl.constexpr,
V_BLOCK_C: tl.constexpr, NUM_SPLITS_C: tl.constexpr,
SPLIT_LEN_C: tl.constexpr, KV_SEQ_LEN_C: tl.constexpr,
):
batch_id = tl.program_id(0)
split_id = tl.program_id(1)
v_block_id = tl.program_id(2)
out_idx = batch_id * NUM_SPLITS_C + split_id
kv_start = batch_id * KV_SEQ_LEN_C + split_id * SPLIT_LEN_C
tl.multiple_of(kv_start, 256)
h_offs = tl.arange(0, NH)
v_start = v_block_id * V_BLOCK_C
v_offs = v_start + tl.arange(0, V_BLOCK_C)
v_mask = v_offs < V_DIM
q_base = batch_id * NH * QK_DIM
tl.multiple_of(q_base, 128)
tl.multiple_of(stride_kv, 32)
tl.assume(stride_kv > 0)
acc = tl.zeros([NH, V_BLOCK_C], dtype=tl.float32)
m_i = tl.full([NH], float("-inf"), dtype=tl.float32)
l_i = tl.zeros([NH], dtype=tl.float32)
score_scale = sm_scale * kv_scale
for kv_offset in range(0, SPLIT_LEN_C, BLOCK_N):
kv_pos = kv_start + kv_offset
tl.multiple_of(kv_pos, 128)
n_offs = tl.arange(0, BLOCK_N)
scores = tl.zeros([NH, BLOCK_N], dtype=tl.float32)
for k_start in range(0, QK_DIM, BLOCK_K):
k_offs = tl.arange(0, BLOCK_K)
k_mask = (k_start + k_offs) < QK_DIM
q_tile = tl.load(Q_ptr + q_base + h_offs[:, None] * QK_DIM + k_start + k_offs[None, :], mask=k_mask[None, :], other=0.0)
k_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + k_start + k_offs[None, :], mask=k_mask[None, :], other=0.0)
scores += tl.dot(q_tile.to(k_tile.dtype), tl.trans(k_tile))
scores *= score_scale
row_max = tl.max(scores, axis=1)
m_new = tl.maximum(m_i, row_max)
alpha = tl.exp(m_i - m_new)
l_i = l_i * alpha
exp_scores = tl.exp(scores - m_new[:, None])
l_i += tl.sum(exp_scores, axis=1)
acc = acc * alpha[:, None]
v_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + v_offs[None, :], mask=v_mask[None, :], other=0.0)
acc += tl.dot(exp_scores.to(v_tile.dtype), v_tile)
m_i = m_new
acc = (acc * kv_scale) / l_i[:, None]
lse_vals = m_i + tl.log(l_i)
tl.store(Mid_O + out_idx * NH * V_DIM + h_offs[:, None] * V_DIM + v_offs[None, :], acc.to(tl.bfloat16), mask=v_mask[None, :])
tl.store(Mid_lse + out_idx * NH + h_offs, lse_vals)
@triton.jit
def _reduce_splitk_parallel_s2(
Mid_O, Mid_lse, O_ptr,
NUM_SPLITS_C: tl.constexpr, V_DIM: tl.constexpr, NH: tl.constexpr, BLOCK_V: tl.constexpr,
):
batch_id = tl.program_id(0)
head_id = tl.program_id(1)
v_block = tl.program_id(2)
v_offs = v_block * BLOCK_V + tl.arange(0, BLOCK_V)
v_mask = v_offs < V_DIM
s_offs = tl.arange(0, NUM_SPLITS_C)
all_lse = tl.load(Mid_lse + (batch_id * NUM_SPLITS_C + s_offs) * NH + head_id)
m_global = tl.max(all_lse, axis=0)
alphas = tl.exp(all_lse - m_global)
l_total = tl.sum(alphas, axis=0)
weights = alphas / l_total
all_partials = tl.load(
Mid_O + (batch_id * NUM_SPLITS_C + s_offs[:, None]) * NH * V_DIM + head_id * V_DIM + v_offs[None, :],
mask=v_mask[None, :], other=0.0,
).to(tl.float32)
result = tl.sum(weights[:, None] * all_partials, axis=0)
out_base = batch_id * NH * V_DIM + head_id * V_DIM
tl.store(O_ptr + out_base + v_offs, result.to(tl.bfloat16), mask=v_mask)
# ============================================================
# S3 triton kernels (batch=32, kv=1024)
# splits=4, SPLIT_LEN=256, V_BLOCK=256, serial reduce w1
# ============================================================
S3_NUM_SPLITS = 4
S3_SPLIT_LEN = 256
S3_KV_SEQ_LEN = 1024
S3_V_BLOCK = 256
@triton.jit
def _flash_decode_fp8_s3_exact_vtile(
Q_ptr, KV_ptr, Mid_O, Mid_lse,
stride_kv: tl.int64, kv_scale,
sm_scale: tl.constexpr, QK_DIM: tl.constexpr, V_DIM: tl.constexpr,
BLOCK_K: tl.constexpr, NH: tl.constexpr, V_BLOCK_C: tl.constexpr,
NUM_SPLITS_C: tl.constexpr, SPLIT_LEN_C: tl.constexpr, KV_SEQ_LEN_C: tl.constexpr,
):
batch_id = tl.program_id(0)
split_id = tl.program_id(1)
v_block_id = tl.program_id(2)
out_idx = batch_id * NUM_SPLITS_C + split_id
kv_pos = batch_id * KV_SEQ_LEN_C + split_id * SPLIT_LEN_C
tl.multiple_of(kv_pos, 128)
h_offs = tl.arange(0, NH)
n_offs = tl.arange(0, SPLIT_LEN_C)
v_start = v_block_id * V_BLOCK_C
v_offs = v_start + tl.arange(0, V_BLOCK_C)
v_mask = v_offs < V_DIM
q_base = batch_id * NH * QK_DIM
tl.multiple_of(q_base, 128)
tl.multiple_of(stride_kv, 32)
tl.assume(stride_kv > 0)
acc = tl.zeros([NH, V_BLOCK_C], dtype=tl.float32)
m_i = tl.full([NH], float("-inf"), dtype=tl.float32)
l_i = tl.zeros([NH], dtype=tl.float32)
score_scale = sm_scale * kv_scale
scores = tl.zeros([NH, SPLIT_LEN_C], dtype=tl.float32)
for k_start in range(0, QK_DIM, BLOCK_K):
k_offs = tl.arange(0, BLOCK_K)
k_mask = (k_start + k_offs) < QK_DIM
q_tile = tl.load(Q_ptr + q_base + h_offs[:, None] * QK_DIM + k_start + k_offs[None, :], mask=k_mask[None, :], other=0.0)
k_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + k_start + k_offs[None, :], mask=k_mask[None, :], other=0.0)
scores += tl.dot(q_tile.to(k_tile.dtype), tl.trans(k_tile))
scores *= score_scale
row_max = tl.max(scores, axis=1)
m_new = tl.maximum(m_i, row_max)
alpha = tl.exp(m_i - m_new)
l_i = l_i * alpha
exp_scores = tl.exp(scores - m_new[:, None])
l_i += tl.sum(exp_scores, axis=1)
acc = acc * alpha[:, None]
v_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + v_offs[None, :], mask=v_mask[None, :], other=0.0)
acc += tl.dot(exp_scores.to(v_tile.dtype), v_tile)
m_i = m_new
acc = (acc * kv_scale) / l_i[:, None]
lse_vals = m_i + tl.log(l_i)
tl.store(Mid_O + out_idx * NH * V_DIM + h_offs[:, None] * V_DIM + v_offs[None, :], acc.to(tl.bfloat16), mask=v_mask[None, :])
tl.store(Mid_lse + out_idx * NH + h_offs, lse_vals)
@triton.jit
def _reduce_splitk_serial_s3(
Mid_O, Mid_lse, O_ptr,
V_DIM: tl.constexpr, NH: tl.constexpr, BLOCK_V: tl.constexpr, NUM_SPLITS_C: tl.constexpr,
):
batch_id = tl.program_id(0)
head_id = tl.program_id(1)
v_block = tl.program_id(2)
v_offs = v_block * BLOCK_V + tl.arange(0, BLOCK_V)
v_mask = v_offs < V_DIM
m_final = tl.full([], float("-inf"), dtype=tl.float32)
l_final = tl.zeros([], dtype=tl.float32)
acc = tl.zeros([BLOCK_V], dtype=tl.float32)
for s in range(NUM_SPLITS_C):
idx = batch_id * NUM_SPLITS_C + s
lse = tl.load(Mid_lse + idx * NH + head_id)
m_new = tl.maximum(m_final, lse)
alpha = tl.exp(m_final - m_new)
beta = tl.exp(lse - m_new)
partial = tl.load(Mid_O + idx * NH * V_DIM + head_id * V_DIM + v_offs, mask=v_mask, other=0.0).to(tl.float32)
acc = acc * alpha + beta * partial
l_final = l_final * alpha + beta
m_final = m_new
acc = acc / l_final
out_base = batch_id * NH * V_DIM + head_id * V_DIM
tl.store(O_ptr + out_base + v_offs, acc.to(tl.bfloat16), mask=v_mask)
# ============================================================
# S5 triton kernels (batch=64, kv=1024)
# splits=4, SPLIT_LEN=256, no V-tiling (full 512), serial reduce w2
# ============================================================
S5_NUM_SPLITS = 4
S5_SPLIT_LEN = 256
S5_KV_SEQ_LEN = 1024
@triton.jit
def _flash_decode_fp8_s5_no_vtile(
Q_ptr, KV_ptr, Mid_O, Mid_lse,
stride_kv: tl.int64, kv_scale,
sm_scale: tl.constexpr, QK_DIM: tl.constexpr, V_DIM: tl.constexpr,
BLOCK_K: tl.constexpr, NH: tl.constexpr,
NUM_SPLITS_C: tl.constexpr, SPLIT_LEN_C: tl.constexpr, KV_SEQ_LEN_C: tl.constexpr,
):
batch_id = tl.program_id(0)
split_id = tl.program_id(1)
out_idx = batch_id * NUM_SPLITS_C + split_id
kv_pos = batch_id * KV_SEQ_LEN_C + split_id * SPLIT_LEN_C
tl.multiple_of(kv_pos, 128)
h_offs = tl.arange(0, NH)
n_offs = tl.arange(0, SPLIT_LEN_C)
v_offs = tl.arange(0, V_DIM)
q_base = batch_id * NH * QK_DIM
tl.multiple_of(q_base, 128)
tl.multiple_of(stride_kv, 32)
tl.assume(stride_kv > 0)
acc = tl.zeros([NH, V_DIM], dtype=tl.float32)
m_i = tl.full([NH], float("-inf"), dtype=tl.float32)
l_i = tl.zeros([NH], dtype=tl.float32)
score_scale = sm_scale * kv_scale
scores = tl.zeros([NH, SPLIT_LEN_C], dtype=tl.float32)
for k_start in range(0, QK_DIM, BLOCK_K):
k_offs = tl.arange(0, BLOCK_K)
k_mask = (k_start + k_offs) < QK_DIM
q_tile = tl.load(Q_ptr + q_base + h_offs[:, None] * QK_DIM + k_start + k_offs[None, :], mask=k_mask[None, :], other=0.0)
k_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + k_start + k_offs[None, :], mask=k_mask[None, :], other=0.0)
scores += tl.dot(q_tile.to(k_tile.dtype), tl.trans(k_tile))
scores *= score_scale
row_max = tl.max(scores, axis=1)
m_new = tl.maximum(m_i, row_max)
alpha = tl.exp(m_i - m_new)
l_i = l_i * alpha
exp_scores = tl.exp(scores - m_new[:, None])
l_i += tl.sum(exp_scores, axis=1)
acc = acc * alpha[:, None]
v_tile = tl.load(KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + v_offs[None, :])
acc += tl.dot(exp_scores.to(v_tile.dtype), v_tile)
m_i = m_new
acc = (acc * kv_scale) / l_i[:, None]
lse_vals = m_i + tl.log(l_i)
tl.store(Mid_O + out_idx * NH * V_DIM + h_offs[:, None] * V_DIM + v_offs[None, :], acc.to(tl.bfloat16))
tl.store(Mid_lse + out_idx * NH + h_offs, lse_vals)
@triton.jit
def _reduce_splitk_serial_s5(
Mid_O, Mid_lse, O_ptr,
V_DIM: tl.constexpr, NH: tl.constexpr, BLOCK_V: tl.constexpr, NUM_SPLITS_C: tl.constexpr,
):
batch_id = tl.program_id(0)
head_id = tl.program_id(1)
v_block = tl.program_id(2)
v_offs = v_block * BLOCK_V + tl.arange(0, BLOCK_V)
v_mask = v_offs < V_DIM
m_final = tl.full([], float("-inf"), dtype=tl.float32)
l_final = tl.zeros([], dtype=tl.float32)
acc = tl.zeros([BLOCK_V], dtype=tl.float32)
for s in range(NUM_SPLITS_C):
idx = batch_id * NUM_SPLITS_C + s
lse = tl.load(Mid_lse + idx * NH + head_id)
m_new = tl.maximum(m_final, lse)
alpha = tl.exp(m_final - m_new)
beta = tl.exp(lse - m_new)
partial = tl.load(Mid_O + idx * NH * V_DIM + head_id * V_DIM + v_offs, mask=v_mask, other=0.0).to(tl.float32)
acc = acc * alpha + beta * partial
l_final = l_final * alpha + beta
m_final = m_new
acc = acc / l_final
out_base = batch_id * NH * V_DIM + head_id * V_DIM
tl.store(O_ptr + out_base + v_offs, acc.to(tl.bfloat16), mask=v_mask)
# ============================================================
# Aiter persistent fp8 path (v143) for S4, S6, S7, S8
# ============================================================
def _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, persistent_splits, fast_mode, kv_gran=64):
key = ("pfp8", batch_size, kv_seq_len, persistent_splits, fast_mode, kv_gran)
if key in _cache:
return _cache[key]
max_q_len = 1
nq, nkv = NUM_HEADS, NUM_KV_HEADS
total_kv = batch_size * kv_seq_len
kv_last_page_len = torch.full((batch_size,), kv_seq_len, dtype=torch.int32, device="cuda")
kv_indices = torch.arange(total_kv, dtype=torch.int32, device="cuda")
info = get_mla_metadata_info_v1(
batch_size, max_q_len, nq, FP8_DTYPE, FP8_DTYPE,
is_sparse=False, fast_mode=fast_mode,
num_kv_splits=persistent_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,
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, kv_gran),
max_seqlen_qo=max_q_len,
uni_seqlen_qo=max_q_len,
fast_mode=fast_mode,
max_split_per_batch=persistent_splits,
intra_batch_mode=True,
dtype_q=FP8_DTYPE,
dtype_kv=FP8_DTYPE,
)
num_partials = reduce_partial_map.size(0)
logits = torch.empty((num_partials, 1, nq, V_HEAD_DIM), dtype=torch.float32, device="cuda")
attn_lse = torch.empty((num_partials, 1, nq, 1), dtype=torch.float32, device="cuda")
o = torch.empty((total_q, nq, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda")
q_fp8 = torch.empty((total_q, nq * QK_HEAD_DIM), dtype=FP8_DTYPE, device="cuda")
q_scale = torch.ones(1, dtype=torch.float32, device="cuda")
_cache[key] = {
"kv_indices": kv_indices, "kv_last_page_len": kv_last_page_len,
"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": logits, "attn_lse": attn_lse, "o": o,
"q_fp8": q_fp8, "q_scale": q_scale,
"num_partials": num_partials,
}
return _cache[key]
# ============================================================
# Triton cache helpers
# ============================================================
def _ensure_triton_cache(shape_key, batch_size, total_q, num_splits, kv_scale_tensor):
key = (shape_key, batch_size, total_q)
if key in _cache:
return _cache[key]
_cache[key] = {
"mid_o": torch.empty((batch_size * num_splits, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
"mid_lse": torch.empty((batch_size * num_splits, NUM_HEADS), dtype=torch.float32, device="cuda"),
"o": torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device="cuda"),
"kv_scale_val": kv_scale_tensor.item(),
}
return _cache[key]
# ============================================================
# Unified dispatch
# ============================================================
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"]
total_q = q.shape[0]
total_kv = batch_size * kv_seq_len
# ---- S1: batch=4, kv=1024 -> triton s1_v45 ----
if batch_size <= 4 and kv_seq_len <= 1024:
kv_fp8, kv_scale = kv_data["fp8"]
kv_flat = kv_fp8.view(total_kv, QK_HEAD_DIM)
c = _ensure_triton_cache("s1_v45", batch_size, total_q, S1_NUM_SPLITS, kv_scale)
grid1 = (batch_size, S1_NUM_SPLITS, triton.cdiv(V_HEAD_DIM, S1_V_BLOCK))
_flash_decode_fp8_s1_exact_vtile[grid1](
q, kv_flat, c["mid_o"], c["mid_lse"],
QK_HEAD_DIM, c["kv_scale_val"], SM_SCALE,
QK_HEAD_DIM, V_HEAD_DIM, 256, NUM_HEADS,
S1_V_BLOCK, S1_NUM_SPLITS, S1_SPLIT_LEN, S1_KV_SEQ_LEN,
num_warps=8, num_stages=3,
)
REDUCE_BLOCK_V = 128
n_v_blocks = triton.cdiv(V_HEAD_DIM, REDUCE_BLOCK_V)
_reduce_splitk_parallel_s1[(batch_size, NUM_HEADS, n_v_blocks)](
c["mid_o"], c["mid_lse"], c["o"],
S1_NUM_SPLITS, V_HEAD_DIM, NUM_HEADS, REDUCE_BLOCK_V,
num_warps=1,
)
return c["o"]
# ---- S2: batch=4, kv=8192 -> triton s2_v57 ----
if batch_size <= 4 and kv_seq_len > 1024:
kv_fp8, kv_scale = kv_data["fp8"]
kv_flat = kv_fp8.view(total_kv, QK_HEAD_DIM)
c = _ensure_triton_cache("s2_v57", batch_size, total_q, S2_NUM_SPLITS, kv_scale)
grid1 = (batch_size, S2_NUM_SPLITS, triton.cdiv(V_HEAD_DIM, S2_V_BLOCK))
_flash_decode_fp8_s2_exact_vtile[grid1](
q, kv_flat, c["mid_o"], c["mid_lse"],
QK_HEAD_DIM, c["kv_scale_val"], SM_SCALE,
QK_HEAD_DIM, V_HEAD_DIM, 256, 256, NUM_HEADS,
S2_V_BLOCK, S2_NUM_SPLITS, S2_SPLIT_LEN, S2_KV_SEQ_LEN,
num_warps=4, num_stages=3,
)
REDUCE_BLOCK_V = 256
n_v_blocks = triton.cdiv(V_HEAD_DIM, REDUCE_BLOCK_V)
_reduce_splitk_parallel_s2[(batch_size, NUM_HEADS, n_v_blocks)](
c["mid_o"], c["mid_lse"], c["o"],
S2_NUM_SPLITS, V_HEAD_DIM, NUM_HEADS, REDUCE_BLOCK_V,
num_warps=2,
)
return c["o"]
# ---- S3: batch=32, kv=1024 -> triton s3_v37 ----
if batch_size <= 32 and kv_seq_len <= 1024:
kv_fp8, kv_scale = kv_data["fp8"]
kv_flat = kv_fp8.view(total_kv, QK_HEAD_DIM)
c = _ensure_triton_cache("s3_v37", batch_size, total_q, S3_NUM_SPLITS, kv_scale)
grid1 = (batch_size, S3_NUM_SPLITS, triton.cdiv(V_HEAD_DIM, S3_V_BLOCK))
_flash_decode_fp8_s3_exact_vtile[grid1](
q, kv_flat, c["mid_o"], c["mid_lse"],
QK_HEAD_DIM, c["kv_scale_val"], SM_SCALE,
QK_HEAD_DIM, V_HEAD_DIM, 256, NUM_HEADS,
S3_V_BLOCK, S3_NUM_SPLITS, S3_SPLIT_LEN, S3_KV_SEQ_LEN,
num_warps=8, num_stages=2,
)
REDUCE_BLOCK_V = 256
n_v_blocks = triton.cdiv(V_HEAD_DIM, REDUCE_BLOCK_V)
_reduce_splitk_serial_s3[(batch_size, NUM_HEADS, n_v_blocks)](
c["mid_o"], c["mid_lse"], c["o"],
V_HEAD_DIM, NUM_HEADS, REDUCE_BLOCK_V, S3_NUM_SPLITS,
num_warps=1,
)
return c["o"]
# ---- S5: batch=64, kv=1024 -> triton s5_v15 ----
if batch_size <= 64 and kv_seq_len <= 1024:
kv_fp8, kv_scale = kv_data["fp8"]
kv_flat = kv_fp8.view(total_kv, QK_HEAD_DIM)
c = _ensure_triton_cache("s5_v15", batch_size, total_q, S5_NUM_SPLITS, kv_scale)
_flash_decode_fp8_s5_no_vtile[(batch_size, S5_NUM_SPLITS)](
q, kv_flat, c["mid_o"], c["mid_lse"],
QK_HEAD_DIM, c["kv_scale_val"], SM_SCALE,
QK_HEAD_DIM, V_HEAD_DIM, 128, NUM_HEADS,
S5_NUM_SPLITS, S5_SPLIT_LEN, S5_KV_SEQ_LEN,
num_warps=8, num_stages=3,
)
reduce_block_v = 512
n_v_blocks = triton.cdiv(V_HEAD_DIM, reduce_block_v)
_reduce_splitk_serial_s5[(batch_size, NUM_HEADS, n_v_blocks)](
c["mid_o"], c["mid_lse"], c["o"],
V_HEAD_DIM, NUM_HEADS, reduce_block_v, S5_NUM_SPLITS,
num_warps=2,
)
return c["o"]
# ---- S4, S6, S7, S8: aiter persistent fp8 (v143) ----
kv_buffer_fp8, kv_scale = kv_data["fp8"]
kv_buffer_4d = kv_buffer_fp8.view(total_kv, PAGE_SIZE, NUM_KV_HEADS, QK_HEAD_DIM)
if total_kv >= 1000000:
# s8 (256, 8192)
splits, fast_mode = 4, False
elif total_kv >= 300000:
# s6 (64, 8192)
splits, fast_mode = 4, False
elif batch_size >= 256:
# s7 (256, 1024)
splits, fast_mode = 4, False
else:
# s4 (32, 8192)
splits, fast_mode = 32, True
kv_gran = 64
c = _ensure_cache_persistent_fp8(batch_size, kv_seq_len, total_q, qo_indptr, kv_indptr, splits, fast_mode, kv_gran)
q_2d = q.view(total_q, NUM_HEADS * QK_HEAD_DIM)
c["q_fp8"].copy_(q_2d)
aiter.mla_decode_stage1_asm_fwd(
c["q_fp8"].view(-1, NUM_HEADS, QK_HEAD_DIM), kv_buffer_4d,
qo_indptr, kv_indptr, c["kv_indices"], c["kv_last_page_len"],
None, c["work_metadata"], c["work_indptr"], c["work_info_set"],
1, PAGE_SIZE, NUM_KV_HEADS, SM_SCALE,
c["logits"], c["attn_lse"], c["o"],
c["q_scale"], kv_scale,
)
aiter.mla_reduce_v1(
c["logits"], c["attn_lse"],
c["reduce_indptr"], c["reduce_final_map"], c["reduce_partial_map"],
1, c["o"], None,
)
return c["o"]
scrolls · 652 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