submission 586148
divc13 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 256 lines, June 9 Researcher Reciprocity License v1.0.
sub75_triton_splitk.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-586148?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:947d94a5963105ca4212c465e88ecb08f29b0175b12c2340a7a5f33b3d8dd700
license declaredunknown
license concludedunknown
authorsdivc13
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
mma
scores += tl.dot(q_chunk, tl.trans(k_chunk))split-k
sub75: Triton flash attention with split-K for CU utilization.tile-m = 16
BLOCK_M = 16Kernel source
sub75_triton_splitk.py256 lines
"""
sub75: Triton flash attention with split-K for CU utilization.
Key improvements over sub49/sub52:
1. Split-K: divide KV sequence across programs for better parallelism
2. Adaptive num_splits: more splits for small batches, fewer for large
3. Torch-based reduction: simple partial result combination
4. Uses bf16 KV (faster than fp8 on Triton/AMD per sub52 results)
Grid: (batch_size, num_m_groups, num_splits)
Each program handles BLOCK_M M-rows and kv_len/num_splits KV tokens.
Outputs partial (acc, max, sumexp) to global memory.
Reduction combines partials with online softmax correction.
"""
import torch
import torch.nn.functional as F
import triton
import triton.language as tl
from task import input_t, output_t
SM_SCALE = 1.0 / (576 ** 0.5)
LOG2E = 1.4426950408889634
@triton.jit
def _flash_splitk(
Q_ptr, KV_ptr,
Acc_ptr, Max_ptr, Sum_ptr,
qo_indptr_ptr, kv_indptr_ptr,
sm_scale_log2e,
stride_q0, stride_q1,
num_heads: tl.constexpr,
num_splits: tl.constexpr,
num_m_groups: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_KV: tl.constexpr,
D_TILE: tl.constexpr,
V_DIM: tl.constexpr,
HEAD_DIM: tl.constexpr,
):
batch = tl.program_id(0)
m_group = tl.program_id(1)
split = tl.program_id(2)
kv_start = tl.load(kv_indptr_ptr + batch)
kv_end = tl.load(kv_indptr_ptr + batch + 1)
kv_len = kv_end - kv_start
q_start = tl.load(qo_indptr_ptr + batch)
q_end = tl.load(qo_indptr_ptr + batch + 1)
q_len = q_end - q_start
total_m = q_len * num_heads
m_start = m_group * BLOCK_M
m_range = tl.arange(0, BLOCK_M)
m_idx = m_start + m_range
m_mask = m_idx < total_m
qi_local = m_idx // num_heads
hi = m_idx % num_heads
qi_global = q_start + qi_local
q_base = qi_global * stride_q0 + hi * stride_q1
# KV range for this split
kv_per_split = (kv_len + num_splits - 1) // num_splits
split_kv_start = split * kv_per_split
split_kv_end = tl.minimum(split_kv_start + kv_per_split, kv_len)
# Online softmax state
m_i = tl.full([BLOCK_M], float('-inf'), dtype=tl.float32)
l_i = tl.zeros([BLOCK_M], dtype=tl.float32)
acc = tl.zeros([BLOCK_M, V_DIM], dtype=tl.float32)
for kv_off in range(split_kv_start, split_kv_end, BLOCK_KV):
kv_range = tl.arange(0, BLOCK_KV)
kv_valid = (kv_off + kv_range) < split_kv_end
kv_base = (kv_start + kv_off + kv_range) * HEAD_DIM
# QK^T tiled by D_TILE
scores = tl.zeros([BLOCK_M, BLOCK_KV], dtype=tl.float32)
for d_off in tl.static_range(0, HEAD_DIM, D_TILE):
d_range = tl.arange(0, D_TILE)
q_chunk = tl.load(
Q_ptr + q_base[:, None] + d_off + d_range[None, :],
mask=m_mask[:, None], other=0.0
).to(tl.bfloat16)
k_chunk = tl.load(
KV_ptr + kv_base[:, None] + d_off + d_range[None, :],
mask=kv_valid[:, None], other=0.0
).to(tl.bfloat16)
scores += tl.dot(q_chunk, tl.trans(k_chunk))
scores *= sm_scale_log2e
scores = tl.where(kv_valid[None, :], scores, float('-inf'))
# Online softmax
m_ij = tl.max(scores, axis=1)
new_m = tl.maximum(m_i, m_ij)
alpha = tl.math.exp2(m_i - new_m)
p = tl.math.exp2(scores - new_m[:, None])
l_i = l_i * alpha + tl.sum(p, axis=1)
acc = acc * alpha[:, None]
m_i = new_m
# OV: V = first 512 dims of KV
v_range = tl.arange(0, V_DIM)
v_block = tl.load(
KV_ptr + kv_base[:, None] + v_range[None, :],
mask=kv_valid[:, None], other=0.0
).to(tl.bfloat16)
acc += tl.dot(p.to(tl.bfloat16), v_block)
# Store partial results
# Layout: flat index = (batch * num_m_groups + m_group) * num_splits + split
flat_idx = (batch * num_m_groups + m_group) * num_splits + split
acc_base = flat_idx * BLOCK_M * V_DIM
ml_base = flat_idx * BLOCK_M
v_range = tl.arange(0, V_DIM)
tl.store(Acc_ptr + acc_base + m_range[:, None] * V_DIM + v_range[None, :],
acc, mask=m_mask[:, None])
tl.store(Max_ptr + ml_base + m_range, m_i, mask=m_mask)
tl.store(Sum_ptr + ml_base + m_range, l_i, mask=m_mask)
@triton.jit
def _reduce_splitk(
Acc_ptr, Max_ptr, Sum_ptr, O_ptr,
qo_indptr_ptr,
stride_o0, stride_o1,
num_heads: tl.constexpr,
num_splits: tl.constexpr,
num_m_groups: tl.constexpr,
BLOCK_M: tl.constexpr,
V_DIM: tl.constexpr,
):
batch = tl.program_id(0)
m_group = tl.program_id(1)
q_start = tl.load(qo_indptr_ptr + batch)
q_end = tl.load(qo_indptr_ptr + batch + 1)
q_len = q_end - q_start
total_m = q_len * num_heads
m_start = m_group * BLOCK_M
m_range = tl.arange(0, BLOCK_M)
m_idx = m_start + m_range
m_mask = m_idx < total_m
qi_local = m_idx // num_heads
hi = m_idx % num_heads
qi_global = q_start + qi_local
# Find global max across splits
base = (batch * num_m_groups + m_group) * num_splits
global_max = tl.full([BLOCK_M], float('-inf'), dtype=tl.float32)
for s in range(num_splits):
flat_idx = base + s
m_s = tl.load(Max_ptr + flat_idx * BLOCK_M + m_range,
mask=m_mask, other=float('-inf'))
global_max = tl.maximum(global_max, m_s)
# Combine partials
v_range = tl.arange(0, V_DIM)
total_acc = tl.zeros([BLOCK_M, V_DIM], dtype=tl.float32)
total_l = tl.zeros([BLOCK_M], dtype=tl.float32)
for s in range(num_splits):
flat_idx = base + s
m_s = tl.load(Max_ptr + flat_idx * BLOCK_M + m_range,
mask=m_mask, other=float('-inf'))
l_s = tl.load(Sum_ptr + flat_idx * BLOCK_M + m_range,
mask=m_mask, other=0.0)
alpha = tl.math.exp2(m_s - global_max)
total_l += l_s * alpha
acc_base = flat_idx * BLOCK_M * V_DIM
acc_s = tl.load(
Acc_ptr + acc_base + m_range[:, None] * V_DIM + v_range[None, :],
mask=m_mask[:, None], other=0.0)
total_acc += acc_s * alpha[:, None]
# Normalize
result = total_acc / total_l[:, None]
# Store final output
o_base = qi_global * stride_o0 + hi * stride_o1
tl.store(O_ptr + o_base[:, None] + v_range[None, :],
result.to(tl.bfloat16), mask=m_mask[:, None])
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
num_heads = config["num_heads"]
v_head_dim = config["v_head_dim"]
q_seq_len = config["q_seq_len"]
batch_size = config["batch_size"]
# Use bf16 KV (faster than fp8 in Triton on AMD)
kv_bf16 = kv_data["bf16"]
kv_flat = kv_bf16.view(-1, 576)
total_q = q.shape[0]
o = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device="cuda")
BLOCK_M = 16
BLOCK_KV = 64
D_TILE = 64
total_m = q_seq_len * num_heads
num_m_groups = (total_m + BLOCK_M - 1) // BLOCK_M
# Adaptive split-K: target ~512 total programs for good CU utilization
total_programs_base = batch_size * num_m_groups
num_splits = max(1, min(32, 512 // max(1, total_programs_base)))
# Allocate partial buffers
total_partials = batch_size * num_m_groups * num_splits
acc_partial = torch.empty((total_partials * BLOCK_M, 512), dtype=torch.float32, device="cuda")
max_partial = torch.empty((total_partials * BLOCK_M,), dtype=torch.float32, device="cuda")
sum_partial = torch.empty((total_partials * BLOCK_M,), dtype=torch.float32, device="cuda")
# Launch flash attention kernel
grid_flash = (batch_size, num_m_groups, num_splits)
_flash_splitk[grid_flash](
q, kv_flat,
acc_partial, max_partial, sum_partial,
qo_indptr, kv_indptr,
SM_SCALE * LOG2E,
q.stride(0), q.stride(1),
num_heads=num_heads,
num_splits=num_splits,
num_m_groups=num_m_groups,
BLOCK_M=BLOCK_M,
BLOCK_KV=BLOCK_KV,
D_TILE=D_TILE,
V_DIM=512,
HEAD_DIM=576,
)
# Launch reduction kernel
grid_reduce = (batch_size, num_m_groups)
_reduce_splitk[grid_reduce](
acc_partial, max_partial, sum_partial, o,
qo_indptr,
o.stride(0), o.stride(1),
num_heads=num_heads,
num_splits=num_splits,
num_m_groups=num_m_groups,
BLOCK_M=BLOCK_M,
V_DIM=512,
)
return o
scrolls · 256 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 585963.
"""- sub74: Pure PyTorch bmm attention with bf16 KV.- No custom kernel, no aiter. Tests batched GEMM approach.+ sub75: Triton flash attention with split-K for CU utilization.- Advantages:- - Single torch.bmm call for QK^T (hipBLAS uses MFMA internally)- - Single torch.bmm call for OV- - No per-batch loops, no kernel launch overhead- - MQA: K/V naturally broadcast via batched matmul+ Key improvements over sub49/sub52:+ 1. Split-K: divide KV sequence across programs for better parallelism+ 2. Adaptive num_splits: more splits for small batches, fewer for large+ 3. Torch-based reduction: simple partial result combination+ 4. Uses bf16 KV (faster than fp8 on Triton/AMD per sub52 results)- Disadvantages:- - bf16 KV = 2x bandwidth of fp8- - Materializes full [bs, qseq*nh, kv_seq] scores tensor+ Grid: (batch_size, num_m_groups, num_splits)+ Each program handles BLOCK_M M-rows and kv_len/num_splits KV tokens.+ Outputs partial (acc, max, sumexp) to global memory.+ Reduction combines partials with online softmax correction."""import torchimport torch.nn.functional as F+ import triton+ import triton.language as tlfrom task import input_t, output_tSM_SCALE = 1.0 / (576 ** 0.5)+ LOG2E = 1.4426950408889634+ @triton.jit+ def _flash_splitk(+ Q_ptr, KV_ptr,+ Acc_ptr, Max_ptr, Sum_ptr,+ qo_indptr_ptr, kv_indptr_ptr,+ sm_scale_log2e,+ stride_q0, stride_q1,+ num_heads: tl.constexpr,+ num_splits: tl.constexpr,+ num_m_groups: tl.constexpr,+ BLOCK_M: tl.constexpr,+ BLOCK_KV: tl.constexpr,+ D_TILE: tl.constexpr,+ V_DIM: tl.constexpr,+ HEAD_DIM: tl.constexpr,+ ):+ batch = tl.program_id(0)+ m_group = tl.program_id(1)+ split = tl.program_id(2)++ kv_start = tl.load(kv_indptr_ptr + batch)+ kv_end = tl.load(kv_indptr_ptr + batch + 1)+ kv_len = kv_end - kv_start+ q_start = tl.load(qo_indptr_ptr + batch)+ q_end = tl.load(qo_indptr_ptr + batch + 1)+ q_len = q_end - q_start++ total_m = q_len * num_heads+ m_start = m_group * BLOCK_M+ m_range = tl.arange(0, BLOCK_M)+ m_idx = m_start + m_range+ m_mask = m_idx < total_m++ qi_local = m_idx // num_heads+ hi = m_idx % num_heads+ qi_global = q_start + qi_local+ q_base = qi_global * stride_q0 + hi * stride_q1++ # KV range for this split+ kv_per_split = (kv_len + num_splits - 1) // num_splits+ split_kv_start = split * kv_per_split+ split_kv_end = tl.minimum(split_kv_start + kv_per_split, kv_len)++ # Online softmax state+ m_i = tl.full([BLOCK_M], float('-inf'), dtype=tl.float32)+ l_i = tl.zeros([BLOCK_M], dtype=tl.float32)+ acc = tl.zeros([BLOCK_M, V_DIM], dtype=tl.float32)++ for kv_off in range(split_kv_start, split_kv_end, BLOCK_KV):+ kv_range = tl.arange(0, BLOCK_KV)+ kv_valid = (kv_off + kv_range) < split_kv_end+ kv_base = (kv_start + kv_off + kv_range) * HEAD_DIM++ # QK^T tiled by D_TILE+ scores = tl.zeros([BLOCK_M, BLOCK_KV], dtype=tl.float32)+ for d_off in tl.static_range(0, HEAD_DIM, D_TILE):+ d_range = tl.arange(0, D_TILE)+ q_chunk = tl.load(+ Q_ptr + q_base[:, None] + d_off + d_range[None, :],+ mask=m_mask[:, None], other=0.0+ ).to(tl.bfloat16)+ k_chunk = tl.load(+ KV_ptr + kv_base[:, None] + d_off + d_range[None, :],+ mask=kv_valid[:, None], other=0.0+ ).to(tl.bfloat16)+ scores += tl.dot(q_chunk, tl.trans(k_chunk))++ scores *= sm_scale_log2e+ scores = tl.where(kv_valid[None, :], scores, float('-inf'))++ # Online softmax+ m_ij = tl.max(scores, axis=1)+ new_m = tl.maximum(m_i, m_ij)+ alpha = tl.math.exp2(m_i - new_m)+ p = tl.math.exp2(scores - new_m[:, None])+ l_i = l_i * alpha + tl.sum(p, axis=1)+ acc = acc * alpha[:, None]+ m_i = new_m++ # OV: V = first 512 dims of KV+ v_range = tl.arange(0, V_DIM)+ v_block = tl.load(+ KV_ptr + kv_base[:, None] + v_range[None, :],+ mask=kv_valid[:, None], other=0.0+ ).to(tl.bfloat16)+ acc += tl.dot(p.to(tl.bfloat16), v_block)++ # Store partial results+ # Layout: flat index = (batch * num_m_groups + m_group) * num_splits + split+ flat_idx = (batch * num_m_groups + m_group) * num_splits + split+ acc_base = flat_idx * BLOCK_M * V_DIM+ ml_base = flat_idx * BLOCK_M++ v_range = tl.arange(0, V_DIM)+ tl.store(Acc_ptr + acc_base + m_range[:, None] * V_DIM + v_range[None, :],+ acc, mask=m_mask[:, None])+ tl.store(Max_ptr + ml_base + m_range, m_i, mask=m_mask)+ tl.store(Sum_ptr + ml_base + m_range, l_i, mask=m_mask)+++ @triton.jit+ def _reduce_splitk(+ Acc_ptr, Max_ptr, Sum_ptr, O_ptr,+ qo_indptr_ptr,+ stride_o0, stride_o1,+ num_heads: tl.constexpr,+ num_splits: tl.constexpr,+ num_m_groups: tl.constexpr,+ BLOCK_M: tl.constexpr,+ V_DIM: tl.constexpr,+ ):+ batch = tl.program_id(0)+ m_group = tl.program_id(1)++ q_start = tl.load(qo_indptr_ptr + batch)+ q_end = tl.load(qo_indptr_ptr + batch + 1)+ q_len = q_end - q_start+ total_m = q_len * num_heads++ m_start = m_group * BLOCK_M+ m_range = tl.arange(0, BLOCK_M)+ m_idx = m_start + m_range+ m_mask = m_idx < total_m++ qi_local = m_idx // num_heads+ hi = m_idx % num_heads+ qi_global = q_start + qi_local++ # Find global max across splits+ base = (batch * num_m_groups + m_group) * num_splits+ global_max = tl.full([BLOCK_M], float('-inf'), dtype=tl.float32)+ for s in range(num_splits):+ flat_idx = base + s+ m_s = tl.load(Max_ptr + flat_idx * BLOCK_M + m_range,+ mask=m_mask, other=float('-inf'))+ global_max = tl.maximum(global_max, m_s)++ # Combine partials+ v_range = tl.arange(0, V_DIM)+ total_acc = tl.zeros([BLOCK_M, V_DIM], dtype=tl.float32)+ total_l = tl.zeros([BLOCK_M], dtype=tl.float32)++ for s in range(num_splits):+ flat_idx = base + s+ m_s = tl.load(Max_ptr + flat_idx * BLOCK_M + m_range,+ mask=m_mask, other=float('-inf'))+ l_s = tl.load(Sum_ptr + flat_idx * BLOCK_M + m_range,+ mask=m_mask, other=0.0)+ alpha = tl.math.exp2(m_s - global_max)+ total_l += l_s * alpha++ acc_base = flat_idx * BLOCK_M * V_DIM+ acc_s = tl.load(+ Acc_ptr + acc_base + m_range[:, None] * V_DIM + v_range[None, :],+ mask=m_mask[:, None], other=0.0)+ total_acc += acc_s * alpha[:, None]++ # Normalize+ result = total_acc / total_l[:, None]++ # Store final output+ o_base = qi_global * stride_o0 + hi * stride_o1+ tl.store(O_ptr + o_base[:, None] + v_range[None, :],+ result.to(tl.bfloat16), mask=m_mask[:, None])++def custom_kernel(data: input_t) -> output_t:q, kv_data, qo_indptr, kv_indptr, config = data- bs = config["batch_size"]- nh = config["num_heads"]- qseq = config["q_seq_len"]- kv_lora_rank = config["kv_lora_rank"]- kv_seq = config["kv_seq_len"]+ num_heads = config["num_heads"]+ v_head_dim = config["v_head_dim"]+ q_seq_len = config["q_seq_len"]+ batch_size = config["batch_size"]- kv_bf16 = kv_data["bf16"] # [total_kv, 1, 576] bf16+ # Use bf16 KV (faster than fp8 in Triton on AMD)+ kv_bf16 = kv_data["bf16"]+ kv_flat = kv_bf16.view(-1, 576)- # Reshape for batched matmul (assumes uniform kv_len across batch)- Q = q.view(bs, qseq, nh, 576).reshape(bs, qseq * nh, 576) # [bs, M, 576]- K = kv_bf16.view(bs, kv_seq, 576) # [bs, N, 576]- V = K[:, :, :kv_lora_rank] # [bs, N, 512]+ total_q = q.shape[0]+ o = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device="cuda")- # QK^T: [bs, M, N] via bf16 GEMM (hipBLAS → MFMA)- scores = torch.bmm(Q, K.transpose(1, 2))- # Softmax in fp32 for numerical stability- scores = F.softmax(scores.float() * SM_SCALE, dim=-1)+ BLOCK_M = 16+ BLOCK_KV = 64+ D_TILE = 64- # OV: [bs, M, 512] via bf16 GEMM- output = torch.bmm(scores.to(torch.bfloat16), V)+ total_m = q_seq_len * num_heads+ num_m_groups = (total_m + BLOCK_M - 1) // BLOCK_M- return output.view(bs, qseq, nh, 512).reshape(bs * qseq, nh, 512)+ # Adaptive split-K: target ~512 total programs for good CU utilization+ total_programs_base = batch_size * num_m_groups+ num_splits = max(1, min(32, 512 // max(1, total_programs_base)))++ # Allocate partial buffers+ total_partials = batch_size * num_m_groups * num_splits+ acc_partial = torch.empty((total_partials * BLOCK_M, 512), dtype=torch.float32, device="cuda")+ max_partial = torch.empty((total_partials * BLOCK_M,), dtype=torch.float32, device="cuda")+ sum_partial = torch.empty((total_partials * BLOCK_M,), dtype=torch.float32, device="cuda")++ # Launch flash attention kernel+ grid_flash = (batch_size, num_m_groups, num_splits)+ _flash_splitk[grid_flash](+ q, kv_flat,+ acc_partial, max_partial, sum_partial,+ qo_indptr, kv_indptr,+ SM_SCALE * LOG2E,+ q.stride(0), q.stride(1),+ num_heads=num_heads,+ num_splits=num_splits,+ num_m_groups=num_m_groups,+ BLOCK_M=BLOCK_M,+ BLOCK_KV=BLOCK_KV,+ D_TILE=D_TILE,+ V_DIM=512,+ HEAD_DIM=576,+ )++ # Launch reduction kernel+ grid_reduce = (batch_size, num_m_groups)+ _reduce_splitk[grid_reduce](+ acc_partial, max_partial, sum_partial, o,+ qo_indptr,+ o.stride(0), o.stride(1),+ num_heads=num_heads,+ num_splits=num_splits,+ num_m_groups=num_m_groups,+ BLOCK_M=BLOCK_M,+ V_DIM=512,+ )++ return o
scrolls · 282 diff lines total
Best evidence level for this revision: reported
JSON