submission 650700
kosox97741 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 387 lines, June 9 Researcher Reciprocity License v1.0.
test_v109_blockk128.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-650700?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:115486f2d0798c94f4c7369070c151b05aa6e218f55b426884ef177a5ede5fd5
license declaredunknown
license concludedunknown
authorskosox97741
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
scores += tl.dot(q_tile, tl.trans(k_tile))num-warps = 8
num_warps=8,online-softmax
m_new = tl.maximum(m_i, row_max)split-k
"""Multi-head split-K flash decode. Grid = (batch, splits)."""stages = 2
num_stages=2,tile-k = 64
Rationale: v102 s8 profile: 9 K-loop iterations per N block with BLOCK_K=64.Kernel source
test_v109_blockk128.py387 lines
"""
test_v109_blockk128: BLOCK_K 64→128 (reduces K-loop from 9 to 5 iterations)
Base: test_v107_block128.py
Direction: CONTINUING from v108 (attempt 12)
Target: ALL shapes — reduce inner K-loop iterations
Change: BLOCK_K 64→128. QK_DIM=576 / 128 = 4.5 → 5 iterations (last one masked).
Reduces K-loop overhead by 44%. Need K-dimension masking for indices >= 576.
Rationale: v102 s8 profile: 9 K-loop iterations per N block with BLOCK_K=64.
LDSBankConflict=24.3%. Fewer K iterations may reduce LDS conflicts.
Scale: INCREMENTAL
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
# MLA constants
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)
_cache = {}
BLOCK_K = 128 # K dim tile (576 / 128 = 5 iterations, last one masked)
# =============================================================================
# Multi-head flash-decode with tl.dot — fp8 KV path
# (Same kernel, BLOCK_N passed as constexpr — Triton JIT caches per-constexpr)
# =============================================================================
@triton.jit
def _flash_decode_multihead_fp8(
Q_ptr, KV_ptr, KV_scale, kv_indptr,
Mid_O, # (batch * splits, NH, V_DIM) fp32
Mid_lse, # (batch * splits, NH) fp32
stride_kv: tl.int64,
sm_scale: tl.constexpr,
QK_DIM: tl.constexpr, # 576
V_DIM: tl.constexpr, # 512
NUM_SPLITS: tl.constexpr,
BLOCK_N: tl.constexpr, # 32 or 64
BLOCK_K: tl.constexpr, # 64
NH: tl.constexpr, # 16
):
"""Multi-head split-K flash decode. Grid = (batch, splits)."""
batch_id = tl.program_id(0)
split_id = tl.program_id(1)
kv_start = tl.load(kv_indptr + batch_id)
kv_end = tl.load(kv_indptr + batch_id + 1)
kv_len = kv_end - kv_start
split_len = tl.cdiv(kv_len, NUM_SPLITS)
my_start = kv_start + split_id * split_len
my_end = tl.minimum(kv_start + (split_id + 1) * split_len, kv_end)
actual_len = my_end - my_start
out_idx = batch_id * NUM_SPLITS + split_id
h_offs = tl.arange(0, NH)
v_offs = tl.arange(0, V_DIM)
if actual_len <= 0:
tl.store(Mid_lse + out_idx * NH + h_offs, tl.full([NH], float("-inf"), dtype=tl.float32))
return
kv_scale_val = tl.load(KV_scale).to(tl.float32)
combined_scale = kv_scale_val * sm_scale
q_base = batch_id * NH * QK_DIM
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)
for kv_offset in range(0, actual_len, BLOCK_N):
n_valid = tl.minimum(BLOCK_N, actual_len - kv_offset)
kv_pos = my_start + kv_offset
n_offs = tl.arange(0, BLOCK_N)
n_mask = n_offs < n_valid
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 # mask for last iteration (576 % 128 = 64)
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,
).to(tl.bfloat16)
k_tile = tl.load(
KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + k_start + k_offs[None, :],
mask=n_mask[:, None] & k_mask[None, :],
other=0.0,
).to(tl.bfloat16)
scores += tl.dot(q_tile, tl.trans(k_tile))
scores *= combined_scale
scores = tl.where(n_mask[None, :], scores, float("-inf"))
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=n_mask[:, None],
other=0.0,
).to(tl.bfloat16)
acc += tl.dot(exp_scores.to(tl.bfloat16), v_tile)
m_i = m_new
acc = (acc * kv_scale_val) / 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,
)
tl.store(Mid_lse + out_idx * NH + h_offs, lse_vals)
# =============================================================================
# Multi-head flash-decode — bf16 KV path
# =============================================================================
@triton.jit
def _flash_decode_multihead_bf16(
Q_ptr, KV_ptr, kv_indptr,
Mid_O, Mid_lse,
stride_kv: tl.int64,
sm_scale: tl.constexpr,
QK_DIM: tl.constexpr,
V_DIM: tl.constexpr,
NUM_SPLITS: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
NH: tl.constexpr,
):
"""Multi-head split-K flash decode with bf16 KV. Grid = (batch, splits)."""
batch_id = tl.program_id(0)
split_id = tl.program_id(1)
kv_start = tl.load(kv_indptr + batch_id)
kv_end = tl.load(kv_indptr + batch_id + 1)
kv_len = kv_end - kv_start
split_len = tl.cdiv(kv_len, NUM_SPLITS)
my_start = kv_start + split_id * split_len
my_end = tl.minimum(kv_start + (split_id + 1) * split_len, kv_end)
actual_len = my_end - my_start
out_idx = batch_id * NUM_SPLITS + split_id
h_offs = tl.arange(0, NH)
v_offs = tl.arange(0, V_DIM)
if actual_len <= 0:
tl.store(Mid_lse + out_idx * NH + h_offs, tl.full([NH], float("-inf"), dtype=tl.float32))
return
q_base = batch_id * NH * QK_DIM
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)
for kv_offset in range(0, actual_len, BLOCK_N):
n_valid = tl.minimum(BLOCK_N, actual_len - kv_offset)
kv_pos = my_start + kv_offset
n_offs = tl.arange(0, BLOCK_N)
n_mask = n_offs < n_valid
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 # mask for last iteration (576 % 128 = 64)
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,
).to(tl.bfloat16)
k_tile = tl.load(
KV_ptr + (kv_pos + n_offs[:, None]) * stride_kv + k_start + k_offs[None, :],
mask=n_mask[:, None] & k_mask[None, :],
other=0.0,
).to(tl.bfloat16)
scores += tl.dot(q_tile, tl.trans(k_tile))
scores *= sm_scale
scores = tl.where(n_mask[None, :], scores, float("-inf"))
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=n_mask[:, None],
other=0.0,
).to(tl.bfloat16)
acc += tl.dot(exp_scores.to(tl.bfloat16), v_tile)
m_i = m_new
acc = acc / 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,
)
tl.store(Mid_lse + out_idx * NH + h_offs, lse_vals)
# =============================================================================
# Split-K reduce kernel
# =============================================================================
@triton.jit
def _reduce_splitk(
Mid_O, # (batch * splits, NH, V_DIM) fp32
Mid_lse, # (batch * splits, NH) fp32
O_ptr, # (total_q, NH, V_DIM) bf16
NUM_SPLITS: tl.constexpr,
V_DIM: tl.constexpr,
NH: tl.constexpr,
BLOCK_V: tl.constexpr,
):
"""Reduce split-K partials. Grid = (batch, NH, cdiv(V_DIM, BLOCK_V))."""
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):
idx = batch_id * NUM_SPLITS + s
lse = tl.load(Mid_lse + idx * NH + head_id)
is_valid = lse > float("-inf")
if is_valid:
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,
)
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)
# =============================================================================
# Cache helper
# =============================================================================
def _ensure_cache(batch_size, kv_seq_len, total_q, num_splits):
key = ("triton_mh_v109", batch_size, kv_seq_len, num_splits)
if key in _cache:
return _cache[key]
_cache[key] = {
"mid_o": torch.empty(
(batch_size * num_splits, NUM_HEADS, V_HEAD_DIM),
dtype=torch.float32, 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",
),
}
return _cache[key]
# =============================================================================
# Main dispatch — fully custom, no aiter
# =============================================================================
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
# Per-shape BLOCK_N: 128 for long KV (halves iterations vs v102's 64), 32 for short
# Profile v102 s8: VGPR=60 (no spill), safe to increase BLOCK_N
if kv_seq_len >= 8192:
block_n = 128
else:
block_n = 32
# Adaptive split count — per-shape tuning (combines best of v100 + v101)
if kv_seq_len <= 1024 and batch_size <= 4:
# s1 (4, 1024): need more splits for CU utilization
# min_tokens=64 → max_splits=16 → programs=64 (matches v100)
min_tokens_per_split = 2 * block_n # 64
elif kv_seq_len <= 1024:
# s3/s5/s7: use min 128 tokens/split (prevents s5 over-splitting)
min_tokens_per_split = 4 * block_n # 128
else:
# Long KV (s2/s4/s6/s8): larger blocks, allow more splits
min_tokens_per_split = 4 * block_n # 256 with block_n=64
max_splits = max(1, kv_seq_len // min_tokens_per_split)
target_programs = 2048
target_splits = max(1, target_programs // batch_size)
num_splits = max(1, min(target_splits, max_splits))
# Cap splits
if num_splits > 64:
num_splits = 64
c = _ensure_cache(batch_size, kv_seq_len, total_q, num_splits)
if batch_size <= 4:
# bf16 path — simpler, no fp8 scale overhead
kv_bf16 = kv_data["bf16"]
kv_flat = kv_bf16.view(total_kv, QK_HEAD_DIM)
grid1 = (batch_size, num_splits)
_flash_decode_multihead_bf16[grid1](
q, kv_flat, kv_indptr,
c["mid_o"], c["mid_lse"],
QK_HEAD_DIM,
SM_SCALE,
QK_HEAD_DIM, V_HEAD_DIM,
num_splits,
block_n, BLOCK_K, NUM_HEADS,
num_warps=8,
num_stages=2,
)
else:
# fp8 path — 2x bandwidth savings
kv_fp8, kv_scale = kv_data["fp8"]
kv_flat = kv_fp8.view(total_kv, QK_HEAD_DIM)
grid1 = (batch_size, num_splits)
_flash_decode_multihead_fp8[grid1](
q, kv_flat, kv_scale, kv_indptr,
c["mid_o"], c["mid_lse"],
QK_HEAD_DIM,
SM_SCALE,
QK_HEAD_DIM, V_HEAD_DIM,
num_splits,
block_n, BLOCK_K, NUM_HEADS,
num_warps=8,
num_stages=2,
)
# Stage 2: reduce split-K partials
REDUCE_BLOCK_V = 128
n_v_blocks = triton.cdiv(V_HEAD_DIM, REDUCE_BLOCK_V)
grid2 = (batch_size, NUM_HEADS, n_v_blocks)
_reduce_splitk[grid2](
c["mid_o"], c["mid_lse"], c["o"],
num_splits, V_HEAD_DIM, NUM_HEADS, REDUCE_BLOCK_V,
num_warps=4,
)
return c["o"]
scrolls · 387 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