submission 667881
migratesky · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 346 lines, June 9 Researcher Reciprocity License v1.0.
v153_prealloc_tuned.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-667881?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:7bcce6a066905200c24a567a1f542f6cc702c690e673056131b1ac82e1ef5d97
license declaredunknown
license concludedunknown
authorsmigratesky
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
scores += tl.dot(q_chunk, tl.trans(kv_chunk.to(tl.bfloat16)))num-warps = 4
num_warps=4,online-softmax
m_new = tl.maximum(m_i, m_ij)stages = 2
num_stages = 2 if kv_seq_len < 4096 else 1tile-n = 64
BLOCK_N = 64 if kv_seq_len >= 4096 else 32Kernel source
v153_prealloc_tuned.py346 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
"""v153: Pre-allocated buffers + per-shape warp tuning.
Base: v146 (geo mean 68.019us).
Changes vs v146:
1. Pre-allocate partial_o, partial_lse, out buffers once (save ~3-5us
per call from avoiding torch.empty).
2. B<=32: use 4 warps (v121 config - marginally better for B32/KV8192).
3. B>=64: keep 8 warps (v146 config - proven best for large batches).
4. BLOCK_V=512 for reduce kernel (single V iteration vs 4 with BLOCK_V=128).
"""
from __future__ import annotations
import sys
import torch
import triton
import triton.language as tl
from task import input_t, output_t
_LOGGED = False
_PARTIAL_O = None
_PARTIAL_LSE = None
_OUT_BUF = None
def _log(msg: str):
print(f"[v153] {msg}", file=sys.stderr, flush=True)
def _ensure_buffers(device):
global _PARTIAL_O, _PARTIAL_LSE, _OUT_BUF
if _PARTIAL_O is None:
_PARTIAL_O = torch.empty(
(2048, 16, 512), dtype=torch.float32, device=device,
)
_PARTIAL_LSE = torch.empty(
(2048, 16), dtype=torch.float32, device=device,
)
_OUT_BUF = torch.empty(
(256, 16, 512), dtype=torch.bfloat16, device=device,
)
@triton.jit
def _flash_decode_stage1(
Q_ptr, KV_ptr, KV_scale_ptr,
Partial_O_ptr, Partial_LSE_ptr,
qo_indptr_ptr, kv_indptr_ptr,
sm_scale,
stride_q_tok, stride_q_h, stride_q_d,
stride_kv_tok, stride_kv_d,
stride_po_row, stride_po_h, stride_po_d,
stride_plse_row, stride_plse_h,
NUM_KV_SPLITS: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_D: tl.constexpr,
NUM_HEADS: tl.constexpr,
QK_DIM: tl.constexpr,
V_DIM: tl.constexpr,
):
batch_id = tl.program_id(0)
split_id = tl.program_id(1)
kv_start = tl.load(kv_indptr_ptr + batch_id)
kv_end = tl.load(kv_indptr_ptr + batch_id + 1)
kv_len = kv_end - kv_start
tokens_per_split = tl.cdiv(kv_len, NUM_KV_SPLITS)
split_start = split_id * tokens_per_split
split_end = tl.minimum(split_start + tokens_per_split, kv_len)
row = batch_id * NUM_KV_SPLITS + split_id
if split_start >= kv_len:
offs_h = tl.arange(0, NUM_HEADS)
tl.store(
Partial_LSE_ptr + row * stride_plse_row
+ offs_h * stride_plse_h,
tl.full([NUM_HEADS], float("-inf"), dtype=tl.float32),
)
return
qo_start = tl.load(qo_indptr_ptr + batch_id)
q_base = Q_ptr + qo_start * stride_q_tok
kv_scale = tl.load(KV_scale_ptr)
LOG2E: tl.constexpr = 1.4426950408889634
qk_scale = sm_scale * kv_scale * LOG2E
m_i = tl.full([NUM_HEADS], float("-inf"), dtype=tl.float32)
l_i = tl.zeros([NUM_HEADS], dtype=tl.float32)
acc = tl.zeros([NUM_HEADS, V_DIM], dtype=tl.float32)
offs_h = tl.arange(0, NUM_HEADS)
offs_n = tl.arange(0, BLOCK_N)
for tok_start in range(split_start, split_end, BLOCK_N):
n_valid = tl.minimum(BLOCK_N, split_end - tok_start)
kv_tok_ids = kv_start + tok_start + offs_n
mask_n = offs_n < n_valid
scores = tl.zeros([NUM_HEADS, BLOCK_N], dtype=tl.float32)
for d_start in range(0, QK_DIM, BLOCK_D):
offs_d = d_start + tl.arange(0, BLOCK_D)
q_chunk = tl.load(
q_base
+ offs_h[:, None] * stride_q_h
+ offs_d[None, :] * stride_q_d,
).to(tl.bfloat16)
kv_chunk = tl.load(
KV_ptr
+ kv_tok_ids[:, None] * stride_kv_tok
+ offs_d[None, :] * stride_kv_d,
mask=mask_n[:, None],
other=0.0,
)
scores += tl.dot(q_chunk, tl.trans(kv_chunk.to(tl.bfloat16)))
scores = scores * qk_scale
scores = tl.where(mask_n[None, :], scores, float("-inf"))
m_ij = tl.max(scores, axis=1)
m_new = tl.maximum(m_i, m_ij)
alpha = tl.math.exp2(m_i - m_new)
p = tl.math.exp2(scores - m_new[:, None])
l_new = alpha * l_i + tl.sum(p, axis=1)
acc = acc * alpha[:, None]
offs_v = tl.arange(0, V_DIM)
v_block = tl.load(
KV_ptr
+ kv_tok_ids[:, None] * stride_kv_tok
+ offs_v[None, :] * stride_kv_d,
mask=mask_n[:, None],
other=0.0,
)
p_bf16 = p.to(tl.bfloat16)
v_bf16 = v_block.to(tl.bfloat16)
acc += tl.dot(p_bf16, v_bf16)
m_i = m_new
l_i = l_new
acc = (acc * kv_scale) / l_i[:, None]
offs_v = tl.arange(0, V_DIM)
tl.store(
Partial_O_ptr + row * stride_po_row
+ offs_h[:, None] * stride_po_h
+ offs_v[None, :] * stride_po_d,
acc.to(tl.float32),
)
lse = tl.math.log2(l_i) + m_i
tl.store(
Partial_LSE_ptr + row * stride_plse_row + offs_h * stride_plse_h,
lse,
)
@triton.jit
def _flash_decode_reduce(
Partial_O_ptr, Partial_LSE_ptr, Out_ptr,
stride_po_row, stride_po_h, stride_po_d,
stride_plse_row, stride_plse_h,
stride_o_tok, stride_o_h, stride_o_d,
qo_indptr_ptr,
NUM_KV_SPLITS: tl.constexpr,
NUM_HEADS: tl.constexpr,
V_DIM: tl.constexpr,
BLOCK_V: tl.constexpr,
):
batch_id = tl.program_id(0)
head_id = tl.program_id(1)
offs_s = tl.arange(0, NUM_KV_SPLITS)
base_row = batch_id * NUM_KV_SPLITS
lse_vals = tl.load(
Partial_LSE_ptr
+ (base_row + offs_s) * stride_plse_row
+ head_id * stride_plse_h,
)
max_lse = tl.max(lse_vals, axis=0)
weights = tl.math.exp2(lse_vals - max_lse)
sum_weights = tl.sum(weights, axis=0)
sum_weights = tl.where(sum_weights > 0.0, sum_weights, 1.0)
qo_start = tl.load(qo_indptr_ptr + batch_id)
offs_v = tl.arange(0, BLOCK_V)
for v_start in range(0, V_DIM, BLOCK_V):
v_offs = v_start + offs_v
mask_v = v_offs < V_DIM
partial_all = tl.load(
Partial_O_ptr
+ (base_row + offs_s)[:, None] * stride_po_row
+ head_id * stride_po_h
+ v_offs[None, :] * stride_po_d,
mask=mask_v[None, :], other=0.0,
)
weighted = partial_all * weights[:, None]
acc = tl.sum(weighted, axis=0) / sum_weights
tl.store(
Out_ptr + qo_start * stride_o_tok
+ head_id * stride_o_h
+ v_offs * stride_o_d,
acc.to(tl.bfloat16), mask=mask_v,
)
def _choose_splits_v121(batch_size: int, kv_seq_len: int) -> int:
CU_COUNT = 256
BLOCK_N_EST = 64 if kv_seq_len >= 4096 else 32
min_tokens_per_split = BLOCK_N_EST * 2
max_splits = max(1, kv_seq_len // min_tokens_per_split)
if batch_size <= 4:
target_programs = CU_COUNT * 16
elif batch_size <= 32:
target_programs = CU_COUNT * 4
elif batch_size <= 64:
target_programs = CU_COUNT * 4
else:
target_programs = CU_COUNT * 8
if batch_size >= target_programs:
return 1
ns = max(1, target_programs // batch_size)
ns = min(ns, max_splits)
ns = 1 << (ns - 1).bit_length() if ns > 1 else 1
return min(ns, 128)
def _choose_splits_v146(batch_size: int, kv_seq_len: int) -> int:
CU_COUNT = 256
BLOCK_N_EST = 64 if kv_seq_len >= 4096 else 32
min_tokens_per_split = BLOCK_N_EST * 2
max_splits = max(1, kv_seq_len // min_tokens_per_split)
if batch_size <= 32:
target_programs = CU_COUNT * 4
elif batch_size <= 64:
target_programs = CU_COUNT * 4
elif batch_size >= 128:
target_programs = CU_COUNT * 4
else:
target_programs = CU_COUNT * 8
if batch_size >= target_programs:
return 1
ns = max(1, target_programs // batch_size)
ns = min(ns, max_splits)
ns = 1 << (ns - 1).bit_length() if ns > 1 else 1
return min(ns, 128)
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
global _LOGGED
q, kv_data, qo_indptr, kv_indptr, config = data
kv_fp8, kv_scale_tensor = kv_data["fp8"]
batch_size = int(config["batch_size"])
kv_seq_len = int(config["kv_seq_len"])
num_heads = int(config["num_heads"])
qk_head_dim = int(config["qk_head_dim"])
v_head_dim = int(config["v_head_dim"])
sm_scale = float(config["sm_scale"])
device = q.device
_ensure_buffers(device)
if batch_size <= 4:
nwarps = 4
BLOCK_N = 64 if kv_seq_len >= 4096 else 32
num_stages = 2 if kv_seq_len < 4096 else 1
num_kv_splits = _choose_splits_v121(batch_size, kv_seq_len)
elif batch_size <= 32:
nwarps = 4
BLOCK_N = 64 if kv_seq_len >= 4096 else 32
num_stages = 2 if kv_seq_len < 4096 else 1
num_kv_splits = _choose_splits_v146(batch_size, kv_seq_len)
else:
nwarps = 8
BLOCK_N = 128 if kv_seq_len >= 4096 else 64
num_stages = 2
num_kv_splits = _choose_splits_v146(batch_size, kv_seq_len)
if not _LOGGED:
_log(
f"B={batch_size} KV={kv_seq_len} splits={num_kv_splits} "
f"warps={nwarps} BN={BLOCK_N} stages={num_stages}"
)
_LOGGED = True
kv_flat = kv_fp8.view(-1, qk_head_dim)
BLOCK_D = 64
n_rows = batch_size * num_kv_splits
partial_o = _PARTIAL_O[:n_rows]
partial_lse = _PARTIAL_LSE[:n_rows]
grid_s1 = (batch_size, num_kv_splits)
_flash_decode_stage1[grid_s1](
q, kv_flat, kv_scale_tensor,
partial_o, partial_lse,
qo_indptr, kv_indptr,
sm_scale,
q.stride(0), q.stride(1), q.stride(2),
kv_flat.stride(0), kv_flat.stride(1),
partial_o.stride(0), partial_o.stride(1), partial_o.stride(2),
partial_lse.stride(0), partial_lse.stride(1),
NUM_KV_SPLITS=num_kv_splits,
BLOCK_N=BLOCK_N,
BLOCK_D=BLOCK_D,
NUM_HEADS=num_heads,
QK_DIM=qk_head_dim,
V_DIM=v_head_dim,
num_warps=nwarps,
num_stages=num_stages,
)
out = _OUT_BUF[:batch_size]
if num_kv_splits == 1:
out.copy_(partial_o[:batch_size].view_as(out).to(torch.bfloat16))
else:
BLOCK_V = 512
grid_r = (batch_size, num_heads)
_flash_decode_reduce[grid_r](
partial_o, partial_lse, out,
partial_o.stride(0), partial_o.stride(1), partial_o.stride(2),
partial_lse.stride(0), partial_lse.stride(1),
out.stride(0), out.stride(1), out.stride(2),
qo_indptr,
NUM_KV_SPLITS=num_kv_splits,
NUM_HEADS=num_heads,
V_DIM=v_head_dim,
BLOCK_V=BLOCK_V,
num_warps=4,
num_stages=1,
)
return out
scrolls · 346 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