submission 747749
divc13 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 253 lines, June 9 Researcher Reciprocity License v1.0.
submission_v145.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-747749?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:88a68723b3bd2639037683032290a3602d4ae08120dfb05f41977617cfcec1d3
license declaredunknown
license concludedunknown
authorsdivc13
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp8
FP8_DTYPE = tl.float8e4nvmma
qk = tl.dot(q_nope, k_nope) + tl.dot(q_rope, k_rope)num-warps = 4
num_warps=4, num_stages=2, waves_per_eu=2,stages = 2
num_warps=4, num_stages=2, waves_per_eu=2,tile-n = 64
NUM_KV_SPLITS=num_splits, BLOCK_N=64, BLOCK_H=16,Kernel source
submission_v145.py253 lines
"""MLA decode Triton v145: v144 refinement — selective NUM_TILES + no .wt.
v144 showed: tiles=1,4 improved (-4 to -6%), tiles=16,64 regressed (+1-2%).
Changes from v144:
1. Only set NUM_TILES constexpr for small tile counts (<=4).
Large tile counts use NUM_TILES=0 (dynamic loop) to avoid
compiler over-scheduling / icache pressure.
2. Remove .wt on mid_acc stores — may hurt large-batch L2 reuse.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
NUM_HEADS = 16
KV_LORA_RANK = 512
QK_HEAD_DIM = 576
V_HEAD_DIM = 512
FP8_DTYPE = tl.float8e4nv
LN2 = tl.constexpr(0.6931471805599453)
_SHAPE_CONFIGS = {
(4, 1024): 16,
(4, 8192): 32,
(32, 1024): 4,
(32, 8192): 8,
(64, 1024): 4,
(64, 8192): 8,
(256, 1024): 1,
(256, 8192): 2,
}
@triton.jit
def _remap_xcd(pid, GRID_SIZE, NUM_XCDS: tl.constexpr = 8):
pids_per_xcd = (GRID_SIZE + NUM_XCDS - 1) // NUM_XCDS
tall_xcds = GRID_SIZE % 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:
return xcd * pids_per_xcd + local_pid
else:
return tall_xcds * pids_per_xcd + (xcd - tall_xcds) * (pids_per_xcd - 1) + local_pid
@triton.jit
def _fast_rcp(x):
return tl.inline_asm_elementwise(
"v_rcp_f32_e32 $0, $1", "=v, v", [x],
dtype=tl.float32, is_pure=True, pack=1,
)
@triton.jit
def _stage1_kernel(
Q, KV, Mid_Acc, Mid_LSE, Out,
qo_indptr, kv_indptr,
sm_scale, kv_scale_ptr,
NUM_KV_SPLITS: tl.constexpr, BLOCK_N: tl.constexpr,
BLOCK_H: tl.constexpr, BLOCK_C: tl.constexpr, BLOCK_R: tl.constexpr,
DIRECT_OUT: tl.constexpr, NUM_TILES: tl.constexpr,
GRID_SIZE: tl.constexpr,
STRIDE_QT: tl.constexpr, STRIDE_QH: tl.constexpr,
STRIDE_KVT: tl.constexpr,
STRIDE_OT: tl.constexpr, STRIDE_OH: tl.constexpr,
):
raw_pid = tl.program_id(0)
pid = _remap_xcd(raw_pid, GRID_SIZE)
cur_batch = pid // NUM_KV_SPLITS
split_id = pid % NUM_KV_SPLITS
heads = tl.arange(0, BLOCK_H)
kv_scale = tl.load(kv_scale_ptr).to(tl.float32)
LOG2E: tl.constexpr = 1.4426950408889634
score_scale_log2 = sm_scale * kv_scale * LOG2E
q_tok = tl.load(qo_indptr + cur_batch)
kv_s = tl.load(kv_indptr + cur_batch)
kv_e = tl.load(kv_indptr + cur_batch + 1)
kv_len = kv_e - kv_s
kv_per_split = tl.cdiv(kv_len, NUM_KV_SPLITS)
sp_start = kv_per_split * split_id
offs_c = tl.arange(0, BLOCK_C)
offs_r = tl.arange(0, BLOCK_R)
q_base = q_tok * STRIDE_QT
q_nope = tl.load(Q + q_base + heads[:, None] * STRIDE_QH + offs_c[None, :]).to(FP8_DTYPE)
q_rope = tl.load(Q + q_base + heads[:, None] * STRIDE_QH + (512 + offs_r[None, :])).to(FP8_DTYPE)
e_max = tl.full([BLOCK_H], value=float("-inf"), dtype=tl.float32)
e_sum = tl.zeros([BLOCK_H], dtype=tl.float32)
acc = tl.zeros([BLOCK_H, BLOCK_C], dtype=tl.float32)
ntiles = NUM_TILES if NUM_TILES > 0 else (tl.minimum(sp_start + kv_per_split, kv_len) - sp_start) // BLOCK_N
for tile_idx in range(ntiles):
start_n = sp_start + tile_idx * BLOCK_N
kv_idx = kv_s + start_n + tl.arange(0, BLOCK_N)
kv_nope = tl.load(KV + kv_idx[:, None] * STRIDE_KVT + offs_c[None, :],
cache_modifier=".cg")
kv_rope = tl.load(KV + kv_idx[:, None] * STRIDE_KVT + (512 + offs_r[None, :]),
cache_modifier=".cg")
k_nope = tl.trans(kv_nope)
k_rope = tl.trans(kv_rope)
qk = tl.dot(q_nope, k_nope) + tl.dot(q_rope, k_rope)
qk = qk * score_scale_log2
new_max = tl.maximum(tl.max(qk, 1), e_max)
rescale = tl.math.exp2(e_max - new_max)
p = tl.math.exp2(qk - new_max[:, None])
acc = acc * rescale[:, None]
acc = acc + tl.dot(p.to(FP8_DTYPE), kv_nope)
e_sum = e_sum * rescale + tl.sum(p, 1)
e_max = new_max
safe_esum = tl.where(e_sum > 0, e_sum, 1.0)
inv_esum = _fast_rcp(safe_esum)
result = acc * inv_esum[:, None]
if DIRECT_OUT:
o_base = q_tok * STRIDE_OT
tl.store(Out + o_base + heads[:, None] * STRIDE_OH + offs_c[None, :],
(result * kv_scale).to(tl.bfloat16))
else:
STRIDE_ACC: tl.constexpr = 512
stride_acc_h = NUM_KV_SPLITS * STRIDE_ACC
stride_acc_b = BLOCK_H * stride_acc_h
acc_base = cur_batch * stride_acc_b + heads * stride_acc_h + split_id * STRIDE_ACC
tl.store(Mid_Acc + acc_base[:, None] + offs_c[None, :], result)
lse = tl.where(e_sum > 0, (e_max + tl.math.log2(e_sum)) * LN2, float("-inf"))
stride_lse_h = NUM_KV_SPLITS
stride_lse_b = BLOCK_H * stride_lse_h
lse_base = cur_batch * stride_lse_b + heads * stride_lse_h + split_id
tl.store(Mid_LSE + lse_base, lse)
@triton.jit
def _stage2_kernel(
Mid_Acc, Mid_LSE, Out, qo_indptr, kv_scale_ptr,
NUM_KV_SPLITS: tl.constexpr, BLOCK_DV: tl.constexpr,
batch: tl.constexpr, GRID_SIZE: tl.constexpr,
STRIDE_OT: tl.constexpr, STRIDE_OH: tl.constexpr,
BLOCK_H: tl.constexpr,
):
LOG2E: tl.constexpr = 1.4426950408889634
raw_pid = tl.program_id(0)
pid = _remap_xcd(raw_pid, GRID_SIZE)
cur_batch = pid % batch
cur_head = pid // batch
kv_scale = tl.load(kv_scale_ptr).to(tl.float32)
q_tok = tl.load(qo_indptr + cur_batch)
offs_d = tl.arange(0, BLOCK_DV)
e_max = float("-inf")
e_sum = 0.0
acc = tl.zeros([BLOCK_DV], dtype=tl.float32)
STRIDE_ACC: tl.constexpr = 512
stride_acc_h = NUM_KV_SPLITS * STRIDE_ACC
stride_acc_b = BLOCK_H * stride_acc_h
acc_base = cur_batch * stride_acc_b + cur_head * stride_acc_h
stride_lse_h = NUM_KV_SPLITS
stride_lse_b = BLOCK_H * stride_lse_h
lse_base = cur_batch * stride_lse_b + cur_head * stride_lse_h
for s in range(NUM_KV_SPLITS):
tv = tl.load(Mid_Acc + acc_base + s * STRIDE_ACC + offs_d)
lse = tl.load(Mid_LSE + lse_base + s)
new_max = tl.maximum(lse, e_max)
old_scale = tl.math.exp2((e_max - new_max) * LOG2E)
exp_lse = tl.math.exp2((lse - new_max) * LOG2E)
acc = acc * old_scale + exp_lse * tv
e_sum = e_sum * old_scale + exp_lse
e_max = new_max
safe_esum = tl.maximum(e_sum, 1e-12)
inv_esum = _fast_rcp(safe_esum)
result = acc * inv_esum * kv_scale
o_base = q_tok * STRIDE_OT + cur_head * STRIDE_OH
tl.store(Out + o_base + offs_d, result.to(tl.bfloat16))
def _get_splits(batch_size, kv_len):
key = (batch_size, kv_len)
if key in _SHAPE_CONFIGS:
return _SHAPE_CONFIGS[key]
max_tiles = max(1, kv_len // 64)
if batch_size >= 256:
return min(2, max_tiles)
elif batch_size >= 64:
return min(max(1, 512 // batch_size), max_tiles)
else:
splits = min(max(1, 768 // batch_size), max_tiles)
while splits > 1 and batch_size * splits > 912:
splits -= 1
return splits
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
kv_fp8, kv_scale_t = kv_data["fp8"]
kv = kv_fp8.view(torch.float8_e4m3fn).view(-1, QK_HEAD_DIM)
batch_size = config["batch_size"]
sm_scale = config["sm_scale"]
total_q = q.shape[0]
kv_len = kv.shape[0] // batch_size if batch_size > 0 else 0
num_splits = _get_splits(batch_size, kv_len)
dev = q.device
direct_out = (num_splits == 1)
kv_per_split = (kv_len + num_splits - 1) // num_splits
num_tiles = kv_per_split // 64
if kv_per_split % 64 != 0 or num_tiles > 4:
num_tiles = 0
if not direct_out:
mid_acc = torch.empty(
(batch_size, NUM_HEADS, num_splits, KV_LORA_RANK),
dtype=torch.float32, device=dev,
)
mid_lse = torch.empty(
(batch_size, NUM_HEADS, num_splits),
dtype=torch.float32, device=dev,
)
else:
mid_acc = torch.empty(1, dtype=torch.float32, device=dev)
mid_lse = torch.empty(1, dtype=torch.float32, device=dev)
out = torch.empty((total_q, NUM_HEADS, V_HEAD_DIM), dtype=torch.bfloat16, device=dev)
grid_size1 = batch_size * num_splits
_stage1_kernel[(grid_size1,)](
q, kv, mid_acc, mid_lse, out, qo_indptr, kv_indptr,
sm_scale, kv_scale_t,
NUM_KV_SPLITS=num_splits, BLOCK_N=64, BLOCK_H=16,
BLOCK_C=512, BLOCK_R=64, DIRECT_OUT=direct_out,
NUM_TILES=num_tiles,
GRID_SIZE=grid_size1,
STRIDE_QT=NUM_HEADS * QK_HEAD_DIM,
STRIDE_QH=QK_HEAD_DIM,
STRIDE_KVT=QK_HEAD_DIM,
STRIDE_OT=NUM_HEADS * V_HEAD_DIM,
STRIDE_OH=V_HEAD_DIM,
num_warps=4, num_stages=2, waves_per_eu=2,
)
if not direct_out:
grid_size2 = NUM_HEADS * batch_size
_stage2_kernel[(grid_size2,)](
mid_acc, mid_lse, out, qo_indptr, kv_scale_t,
NUM_KV_SPLITS=num_splits, BLOCK_DV=512,
batch=batch_size, GRID_SIZE=grid_size2,
STRIDE_OT=NUM_HEADS * V_HEAD_DIM,
STRIDE_OH=V_HEAD_DIM,
BLOCK_H=16,
num_warps=4, num_stages=2,
)
return out
scrolls · 253 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 746709.
- """MLA decode Triton v140: v135 + separated LSE tensor (stride=512).+ """MLA decode Triton v145: v144 refinement — selective NUM_TILES + no .wt.- Isolates the separated tensor layout from v138 (which also changed- stage2 nw=2). Keeps stage2 nw=4 to avoid the bs32/kv8k regression.-- mid_acc: [batch, heads, splits, 512] stride=512 (power of 2, shift)- mid_lse: [batch, heads, splits] separate small tensor+ v144 showed: tiles=1,4 improved (-4 to -6%), tiles=16,64 regressed (+1-2%).+ Changes from v144:+ 1. Only set NUM_TILES constexpr for small tile counts (<=4).+ Large tile counts use NUM_TILES=0 (dynamic loop) to avoid+ compiler over-scheduling / icache pressure.+ 2. Remove .wt on mid_acc stores — may hurt large-batch L2 reuse."""import torch⋯ 6 unchanged linesQK_HEAD_DIM = 576V_HEAD_DIM = 512FP8_DTYPE = tl.float8e4nv- LOG2E = tl.constexpr(1.4426950408889634)LN2 = tl.constexpr(0.6931471805599453)_SHAPE_CONFIGS = {⋯ 36 unchanged linessm_scale, kv_scale_ptr,NUM_KV_SPLITS: tl.constexpr, BLOCK_N: tl.constexpr,BLOCK_H: tl.constexpr, BLOCK_C: tl.constexpr, BLOCK_R: tl.constexpr,- DIRECT_OUT: tl.constexpr, GRID_SIZE: tl.constexpr,+ DIRECT_OUT: tl.constexpr, NUM_TILES: tl.constexpr,+ GRID_SIZE: tl.constexpr,STRIDE_QT: tl.constexpr, STRIDE_QH: tl.constexpr,STRIDE_KVT: tl.constexpr,STRIDE_OT: tl.constexpr, STRIDE_OH: tl.constexpr,⋯ 4 unchanged linessplit_id = pid % NUM_KV_SPLITSheads = tl.arange(0, BLOCK_H)kv_scale = tl.load(kv_scale_ptr).to(tl.float32)- score_scale = sm_scale * kv_scale+ LOG2E: tl.constexpr = 1.4426950408889634+ score_scale_log2 = sm_scale * kv_scale * LOG2Eq_tok = tl.load(qo_indptr + cur_batch)kv_s = tl.load(kv_indptr + cur_batch)kv_e = tl.load(kv_indptr + cur_batch + 1)kv_len = kv_e - kv_skv_per_split = tl.cdiv(kv_len, NUM_KV_SPLITS)sp_start = kv_per_split * split_id- sp_end = tl.minimum(sp_start + kv_per_split, kv_len)offs_c = tl.arange(0, BLOCK_C)offs_r = tl.arange(0, BLOCK_R)q_base = q_tok * STRIDE_QT⋯ 2 unchanged linese_max = tl.full([BLOCK_H], value=float("-inf"), dtype=tl.float32)e_sum = tl.zeros([BLOCK_H], dtype=tl.float32)acc = tl.zeros([BLOCK_H, BLOCK_C], dtype=tl.float32)- num_tokens = sp_end - sp_start- num_full = num_tokens // BLOCK_N- full_end = sp_start + num_full * BLOCK_N- for start_n in range(sp_start, full_end, BLOCK_N):+ ntiles = NUM_TILES if NUM_TILES > 0 else (tl.minimum(sp_start + kv_per_split, kv_len) - sp_start) // BLOCK_N+ for tile_idx in range(ntiles):+ start_n = sp_start + tile_idx * BLOCK_Nkv_idx = kv_s + start_n + tl.arange(0, BLOCK_N)kv_nope = tl.load(KV + kv_idx[:, None] * STRIDE_KVT + offs_c[None, :],cache_modifier=".cg")⋯ 2 unchanged linesk_nope = tl.trans(kv_nope)k_rope = tl.trans(kv_rope)qk = tl.dot(q_nope, k_nope) + tl.dot(q_rope, k_rope)- qk = qk * score_scale+ qk = qk * score_scale_log2new_max = tl.maximum(tl.max(qk, 1), e_max)- rescale = tl.math.exp2((e_max - new_max) * LOG2E)- p = tl.math.exp2((qk - new_max[:, None]) * LOG2E)+ rescale = tl.math.exp2(e_max - new_max)+ p = tl.math.exp2(qk - new_max[:, None])acc = acc * rescale[:, None]acc = acc + tl.dot(p.to(FP8_DTYPE), kv_nope)e_sum = e_sum * rescale + tl.sum(p, 1)e_max = new_max- if full_end < sp_end:- offs_n = full_end + tl.arange(0, BLOCK_N)- mask_n = offs_n < sp_end- kv_idx = kv_s + offs_n- kv_nope = tl.load(KV + kv_idx[:, None] * STRIDE_KVT + offs_c[None, :],- mask=mask_n[:, None], other=0.0, cache_modifier=".cg")- kv_rope = tl.load(KV + kv_idx[:, None] * STRIDE_KVT + (512 + offs_r[None, :]),- mask=mask_n[:, None], other=0.0, cache_modifier=".cg")- k_nope = tl.trans(kv_nope)- k_rope = tl.trans(kv_rope)- qk = tl.dot(q_nope, k_nope) + tl.dot(q_rope, k_rope)- qk = qk * score_scale- qk = tl.where(mask_n[None, :], qk, float("-inf"))- new_max = tl.maximum(tl.max(qk, 1), e_max)- rescale = tl.math.exp2((e_max - new_max) * LOG2E)- p = tl.math.exp2((qk - new_max[:, None]) * LOG2E)- acc = acc * rescale[:, None]- acc = acc + tl.dot(p.to(FP8_DTYPE), kv_nope)- e_sum = e_sum * rescale + tl.sum(p, 1)- e_max = new_max-safe_esum = tl.where(e_sum > 0, e_sum, 1.0)inv_esum = _fast_rcp(safe_esum)result = acc * inv_esum[:, None]⋯ 7 unchanged linesstride_acc_b = BLOCK_H * stride_acc_hacc_base = cur_batch * stride_acc_b + heads * stride_acc_h + split_id * STRIDE_ACCtl.store(Mid_Acc + acc_base[:, None] + offs_c[None, :], result)- lse = tl.where(e_sum > 0, e_max + tl.math.log2(e_sum) * LN2, float("-inf"))+ lse = tl.where(e_sum > 0, (e_max + tl.math.log2(e_sum)) * LN2, float("-inf"))stride_lse_h = NUM_KV_SPLITSstride_lse_b = BLOCK_H * stride_lse_hlse_base = cur_batch * stride_lse_b + heads * stride_lse_h + split_id⋯ 8 unchanged linesSTRIDE_OT: tl.constexpr, STRIDE_OH: tl.constexpr,BLOCK_H: tl.constexpr,):+ LOG2E: tl.constexpr = 1.4426950408889634raw_pid = tl.program_id(0)pid = _remap_xcd(raw_pid, GRID_SIZE)cur_batch = pid % batch⋯ 57 unchanged linesdev = q.devicedirect_out = (num_splits == 1)+ kv_per_split = (kv_len + num_splits - 1) // num_splits+ num_tiles = kv_per_split // 64+ if kv_per_split % 64 != 0 or num_tiles > 4:+ num_tiles = 0+if not direct_out:mid_acc = torch.empty((batch_size, NUM_HEADS, num_splits, KV_LORA_RANK),⋯ 15 unchanged linessm_scale, kv_scale_t,NUM_KV_SPLITS=num_splits, BLOCK_N=64, BLOCK_H=16,BLOCK_C=512, BLOCK_R=64, DIRECT_OUT=direct_out,+ NUM_TILES=num_tiles,GRID_SIZE=grid_size1,STRIDE_QT=NUM_HEADS * QK_HEAD_DIM,STRIDE_QH=QK_HEAD_DIM,
scrolls · 144 diff lines total
Best evidence level for this revision: reported
JSON