submission 588924
divc13 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 358 lines, June 9 Researcher Reciprocity License v1.0.
sub120_hybrid.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-588924?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:df0d5b2279f1197ba978d9202b82340d5ad49219489c2d5ec7d5b884e708bb39
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
def _flash_splitk(tile-m = 16
BLOCK_M = 16 if total_m <= 16 else (32 if total_m <= 32 else 64)Kernel source
sub120_hybrid.py358 lines
"""
sub120: Hybrid kernel — dispatches between batched bmm and Triton flash attention.
Key insight from benchmarks:
- bmm is VERY fast for small kv_len (bs=4,kv=1024: 23.6µs vs Triton ~35µs)
- bmm is VERY slow for large kv_len (bs=128,kv=8192,qseq=4: 976µs vs Triton ~574µs)
Strategy:
- If kv_len * qseq <= threshold: use batched bmm (avoids kernel launch overhead)
- Else: use Triton flash attention (avoids materializing full score matrix)
Also uses fp8 KV for both paths where possible.
"""
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
SM_SCALE_LOG2E = SM_SCALE * LOG2E
# ==================== Triton Flash Attention (from sub111) ====================
@triton.jit
def _flash_fused(
Q_ptr, KV_ptr, O_ptr,
qo_indptr_ptr, kv_indptr_ptr,
sm_scale_log2e,
stride_q0, stride_q1,
stride_kv0,
stride_o0, stride_o1,
num_heads: 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)
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
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(0, kv_len, BLOCK_KV):
kv_range = tl.arange(0, BLOCK_KV)
kv_valid = (kv_off + kv_range) < kv_len
kv_base = (kv_start + kv_off + kv_range) * stride_kv0
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'))
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
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)
result = acc / l_i[:, None]
o_base = qi_global * stride_o0 + hi * stride_o1
v_range = tl.arange(0, V_DIM)
tl.store(O_ptr + o_base[:, None] + v_range[None, :],
result.to(tl.bfloat16), mask=m_mask[:, None])
@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,
stride_kv0,
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_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)
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) * stride_kv0
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'))
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
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)
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
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):
m_s = tl.load(Max_ptr + (base + s) * BLOCK_M + m_range, mask=m_mask, other=float('-inf'))
global_max = tl.maximum(global_max, m_s)
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]
result = total_acc / total_l[:, None]
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])
# ==================== BMM Path ====================
def _bmm_attention(q, kv_bf16, qo_indptr, kv_indptr, config):
num_heads = config["num_heads"]
v_head_dim = config["v_head_dim"]
batch_size = config["batch_size"]
q_seq_len = config["q_seq_len"]
kv_seq_len = config["kv_seq_len"]
total_q = q.shape[0]
q_batched = q.view(batch_size, q_seq_len, num_heads, 576).reshape(batch_size, q_seq_len * num_heads, 576)
kv_batched = kv_bf16.view(batch_size, kv_seq_len, 576)
scores = torch.bmm(q_batched, kv_batched.transpose(1, 2))
scores.mul_(SM_SCALE)
scores = F.softmax(scores, dim=-1)
v_batched = kv_batched[:, :, :v_head_dim]
output = torch.bmm(scores.to(v_batched.dtype), v_batched)
return output.view(batch_size, q_seq_len, num_heads, v_head_dim).reshape(total_q, num_heads, v_head_dim).to(torch.bfloat16)
# ==================== Triton Path ====================
def _triton_flash_attention(q, kv_flat, qo_indptr, kv_indptr, config):
num_heads = config["num_heads"]
v_head_dim = config["v_head_dim"]
q_seq_len = config["q_seq_len"]
batch_size = config["batch_size"]
total_q = q.shape[0]
o = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device="cuda")
total_m = q_seq_len * num_heads
BLOCK_M = 16 if total_m <= 16 else (32 if total_m <= 32 else 64)
BLOCK_KV = 64
D_TILE = 64
num_m_groups = (total_m + BLOCK_M - 1) // BLOCK_M
total_programs_base = batch_size * num_m_groups
if total_programs_base >= 128:
grid = (batch_size, num_m_groups)
_flash_fused[grid](
q, kv_flat, o, qo_indptr, kv_indptr,
SM_SCALE_LOG2E,
q.stride(0), q.stride(1), kv_flat.stride(0),
o.stride(0), o.stride(1),
num_heads=num_heads, BLOCK_M=BLOCK_M, BLOCK_KV=BLOCK_KV,
D_TILE=D_TILE, V_DIM=512, HEAD_DIM=576,
)
else:
num_splits = max(1, min(32, 512 // max(1, total_programs_base)))
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")
_flash_splitk[(batch_size, num_m_groups, num_splits)](
q, kv_flat, acc_partial, max_partial, sum_partial,
qo_indptr, kv_indptr, SM_SCALE_LOG2E,
q.stride(0), q.stride(1), kv_flat.stride(0),
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,
)
_reduce_splitk[(batch_size, num_m_groups)](
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
# ==================== Dispatch ====================
def custom_kernel(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
kv_seq_len = config["kv_seq_len"]
q_seq_len = config["q_seq_len"]
batch_size = config["batch_size"]
num_heads = config["num_heads"]
# Heuristic: bmm is better when the score matrix is small
# score_matrix_size = batch_size * q_seq_len * num_heads * kv_seq_len
# bmm materializes the full score matrix in memory
# Flash attention doesn't, so it wins for large score matrices
score_size = q_seq_len * kv_seq_len
if score_size <= 4096: # e.g., qseq=1, kv≤4096 or qseq=4, kv≤1024
# Use batched bmm — faster for small problems
kv_bf16 = kv_data["bf16"]
return _bmm_attention(q, kv_bf16, qo_indptr, kv_indptr, config)
else:
# Use Triton flash attention — better for large score matrices
kv_bf16 = kv_data["bf16"]
kv_flat = kv_bf16.view(-1, 576)
return _triton_flash_attention(q, kv_flat, qo_indptr, kv_indptr, config)
scrolls · 358 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 586148.
"""- sub75: Triton flash attention with split-K for CU utilization.+ sub120: Hybrid kernel — dispatches between batched bmm and Triton flash attention.- 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)+ Key insight from benchmarks:+ - bmm is VERY fast for small kv_len (bs=4,kv=1024: 23.6µs vs Triton ~35µs)+ - bmm is VERY slow for large kv_len (bs=128,kv=8192,qseq=4: 976µs vs Triton ~574µs)- 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.+ Strategy:+ - If kv_len * qseq <= threshold: use batched bmm (avoids kernel launch overhead)+ - Else: use Triton flash attention (avoids materializing full score matrix)++ Also uses fp8 KV for both paths where possible."""import torch⋯ 4 unchanged linesSM_SCALE = 1.0 / (576 ** 0.5)LOG2E = 1.4426950408889634+ SM_SCALE_LOG2E = SM_SCALE * LOG2E+ # ==================== Triton Flash Attention (from sub111) ====================+@triton.jit+ def _flash_fused(+ Q_ptr, KV_ptr, O_ptr,+ qo_indptr_ptr, kv_indptr_ptr,+ sm_scale_log2e,+ stride_q0, stride_q1,+ stride_kv0,+ stride_o0, stride_o1,+ num_heads: 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)++ 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++ 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(0, kv_len, BLOCK_KV):+ kv_range = tl.arange(0, BLOCK_KV)+ kv_valid = (kv_off + kv_range) < kv_len+ kv_base = (kv_start + kv_off + kv_range) * stride_kv0++ 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'))++ 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++ 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)++ result = acc / l_i[:, None]+ o_base = qi_global * stride_o0 + hi * stride_o1+ v_range = tl.arange(0, V_DIM)+ tl.store(O_ptr + o_base[:, None] + v_range[None, :],+ result.to(tl.bfloat16), mask=m_mask[:, None])+++ @triton.jitdef _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,+ stride_kv0,num_heads: tl.constexpr,num_splits: tl.constexpr,num_m_groups: tl.constexpr,⋯ 25 unchanged linesqi_global = q_start + qi_localq_base = qi_global * stride_q0 + hi * stride_q1- # KV range for this splitkv_per_split = (kv_len + num_splits - 1) // num_splitssplit_kv_start = split * kv_per_splitsplit_kv_end = tl.minimum(split_kv_start + kv_per_split, kv_len)- # Online softmax statem_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)⋯ 1 unchanged linesfor 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+ kv_base = (kv_start + kv_off + kv_range) * stride_kv0- # QK^T tiled by D_TILEscores = 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)⋯ 10 unchanged linesscores *= sm_scale_log2escores = tl.where(kv_valid[None, :], scores, float('-inf'))- # Online softmaxm_ij = tl.max(scores, axis=1)new_m = tl.maximum(m_i, m_ij)alpha = tl.math.exp2(m_i - new_m)⋯ 2 unchanged linesacc = acc * alpha[:, None]m_i = new_m- # OV: V = first 512 dims of KVv_range = tl.arange(0, V_DIM)v_block = tl.load(KV_ptr + kv_base[:, None] + v_range[None, :],⋯ 1 unchanged lines).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 + splitflat_idx = (batch * num_m_groups + m_group) * num_splits + splitacc_base = flat_idx * BLOCK_M * V_DIMml_base = flat_idx * BLOCK_M⋯ 18 unchanged lines):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⋯ 3 unchanged linesm_range = tl.arange(0, BLOCK_M)m_idx = m_start + m_rangem_mask = m_idx < total_m-qi_local = m_idx // num_headshi = m_idx % num_headsqi_global = q_start + qi_local- # Find global max across splitsbase = (batch * num_m_groups + m_group) * num_splitsglobal_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'))+ m_s = tl.load(Max_ptr + (base + s) * BLOCK_M + m_range, mask=m_mask, other=float('-inf'))global_max = tl.maximum(global_max, m_s)- # Combine partialsv_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)+ 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)+ 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]- # Normalizeresult = total_acc / total_l[:, None]-- # Store final outputo_base = qi_global * stride_o0 + hi * stride_o1tl.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+ # ==================== BMM Path ====================+ def _bmm_attention(q, kv_bf16, qo_indptr, kv_indptr, config):num_heads = config["num_heads"]v_head_dim = config["v_head_dim"]- q_seq_len = config["q_seq_len"]batch_size = config["batch_size"]+ q_seq_len = config["q_seq_len"]+ kv_seq_len = config["kv_seq_len"]+ total_q = q.shape[0]- # Use bf16 KV (faster than fp8 in Triton on AMD)- kv_bf16 = kv_data["bf16"]- kv_flat = kv_bf16.view(-1, 576)+ q_batched = q.view(batch_size, q_seq_len, num_heads, 576).reshape(batch_size, q_seq_len * num_heads, 576)+ kv_batched = kv_bf16.view(batch_size, kv_seq_len, 576)+ scores = torch.bmm(q_batched, kv_batched.transpose(1, 2))+ scores.mul_(SM_SCALE)+ scores = F.softmax(scores, dim=-1)++ v_batched = kv_batched[:, :, :v_head_dim]+ output = torch.bmm(scores.to(v_batched.dtype), v_batched)++ return output.view(batch_size, q_seq_len, num_heads, v_head_dim).reshape(total_q, num_heads, v_head_dim).to(torch.bfloat16)+++ # ==================== Triton Path ====================++ def _triton_flash_attention(q, kv_flat, qo_indptr, kv_indptr, config):+ num_heads = config["num_heads"]+ v_head_dim = config["v_head_dim"]+ q_seq_len = config["q_seq_len"]+ batch_size = config["batch_size"]+total_q = q.shape[0]o = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device="cuda")- BLOCK_M = 16+ total_m = q_seq_len * num_heads+ BLOCK_M = 16 if total_m <= 16 else (32 if total_m <= 32 else 64)BLOCK_KV = 64D_TILE = 64- total_m = q_seq_len * num_headsnum_m_groups = (total_m + BLOCK_M - 1) // BLOCK_M-- # Adaptive split-K: target ~512 total programs for good CU utilizationtotal_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")+ if total_programs_base >= 128:+ grid = (batch_size, num_m_groups)+ _flash_fused[grid](+ q, kv_flat, o, qo_indptr, kv_indptr,+ SM_SCALE_LOG2E,+ q.stride(0), q.stride(1), kv_flat.stride(0),+ o.stride(0), o.stride(1),+ num_heads=num_heads, BLOCK_M=BLOCK_M, BLOCK_KV=BLOCK_KV,+ D_TILE=D_TILE, V_DIM=512, HEAD_DIM=576,+ )+ else:+ num_splits = max(1, min(32, 512 // max(1, total_programs_base)))+ 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,- )+ _flash_splitk[(batch_size, num_m_groups, num_splits)](+ q, kv_flat, acc_partial, max_partial, sum_partial,+ qo_indptr, kv_indptr, SM_SCALE_LOG2E,+ q.stride(0), q.stride(1), kv_flat.stride(0),+ 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,+ )+ _reduce_splitk[(batch_size, num_m_groups)](+ 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,+ )- # 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+++ # ==================== Dispatch ====================++ def custom_kernel(data: input_t) -> output_t:+ q, kv_data, qo_indptr, kv_indptr, config = data++ kv_seq_len = config["kv_seq_len"]+ q_seq_len = config["q_seq_len"]+ batch_size = config["batch_size"]+ num_heads = config["num_heads"]++ # Heuristic: bmm is better when the score matrix is small+ # score_matrix_size = batch_size * q_seq_len * num_heads * kv_seq_len+ # bmm materializes the full score matrix in memory+ # Flash attention doesn't, so it wins for large score matrices+ score_size = q_seq_len * kv_seq_len++ if score_size <= 4096: # e.g., qseq=1, kv≤4096 or qseq=4, kv≤1024+ # Use batched bmm — faster for small problems+ kv_bf16 = kv_data["bf16"]+ return _bmm_attention(q, kv_bf16, qo_indptr, kv_indptr, config)+ else:+ # Use Triton flash attention — better for large score matrices+ kv_bf16 = kv_data["bf16"]+ kv_flat = kv_bf16.view(-1, 576)+ return _triton_flash_attention(q, kv_flat, qo_indptr, kv_indptr, config)
scrolls · 384 diff lines total
Best evidence level for this revision: reported
JSON