submission 688800
SSS · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 939 lines, June 9 Researcher Reciprocity License v1.0.
vsota.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mixed-mla-688800?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:07e93b5c535a52737f4863379c1142cec61fbb010ceafa965e88e36605c9c411
license declaredunknown
license concludedunknown
authorsSSS
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
@triton.autotune(mma
qk = tl.dot(q_lora_fp8, tl.trans(kv_lora_fp8))num-warps = 4
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64}, num_stages=2, num_warps=4),online-softmax
m_i_new = tl.maximum(m_i, m_ij)split-k
use_split_kv = base_blocks < target_blocksstages = 2
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64}, num_stages=2, num_warps=4),Kernel source
vsota.py939 lines
#!POPCORN leaderboard amd-mixed-mla
#!POPCORN gpu MI355X
# vsota
"""
```
seed: 4217; qseqlen: 1; kvseqlen: 1024; batchsize: 4
⏱ 20.7 ± 0.02 µs
⚡ 19.9 µs 🐌 25.3 µs
seed: 4220; qseqlen: 1; kvseqlen: 8192; batchsize: 4
⏱ 35.1 ± 0.04 µs
⚡ 34.3 µs 🐌 39.2 µs
seed: 5412; qseqlen: 1; kvseqlen: 1024; batchsize: 32
⏱ 22.6 ± 0.02 µs
⚡ 21.8 µs 🐌 28.0 µs
seed: 5415; qseqlen: 1; kvseqlen: 8192; batchsize: 32
⏱ 65.3 ± 0.07 µs
⚡ 64.2 µs 🐌 68.9 µs
seed: 1357; qseqlen: 1; kvseqlen: 1024; batchsize: 64
⏱ 27.7 ± 0.03 µs
⚡ 26.8 µs 🐌 32.4 µs
seed: 1360; qseqlen: 1; kvseqlen: 8192; batchsize: 64
⏱ 113 ± 0.1 µs
⚡ 111 µs 🐌 116 µs
seed: 9823; qseqlen: 1; kvseqlen: 1024; batchsize: 256
⏱ 60.2 ± 0.06 µs
⚡ 59.0 µs 🐌 63.3 µs
seed: 9826; qseqlen: 1; kvseqlen: 8192; batchsize: 256
⏱ 303 ± 0.3 µs
⚡ 294 µs 🐌 316 µs
```
```
"""
"""
## Benchmarks:
# To be filled by benchmark run
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
from aiter import dtypes as aiter_dtypes
FP8_DTYPE = aiter_dtypes.fp8
SHORT_KV_MAX = 1024
SHORT_KV_TARGET_BLOCKS_SMALL = 128
SHORT_KV_TARGET_BLOCKS = 256
SHORT_KV_MAX_SPLITS = 32
SHORT_KV_SMALL_REDUCE_MAX_TOTAL_TH = 64
SHORT_KV_SMALL_REDUCE_BLOCK_D = 32
def custom_kernel(data: input_t) -> output_t:
return custom_kernel_v41(data)
# ==============================================================================
# No-Split Kernel
# ==============================================================================
@triton.autotune(
configs=[
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64}, num_stages=2, num_warps=4),
triton.Config({'BLOCK_M': 32, 'BLOCK_N': 64}, num_stages=2, num_warps=4),
triton.Config({'BLOCK_M': 64, 'BLOCK_N': 64}, num_stages=2, num_warps=8),
triton.Config({'BLOCK_M': 32, 'BLOCK_N': 128}, num_stages=2, num_warps=8),
triton.Config({'BLOCK_M': 64, 'BLOCK_N': 128}, num_stages=2, num_warps=8),
],
key=['KV_LORA_RANK', 'QK_ROPE_DIM'],
)
@triton.jit
def _mla_v11_no_split(
Q, KV, Out, qo_indptr, kv_indptr, kv_scale_ptr, sm_scale,
stride_q_t, stride_q_h, stride_q_d,
stride_kv_t, stride_kv_d,
stride_o_t, stride_o_h, stride_o_d,
num_heads: tl.constexpr, KV_LORA_RANK: tl.constexpr, QK_ROPE_DIM: tl.constexpr, v_head_dim: tl.constexpr,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
batch_idx = tl.program_id(0)
m_block_idx = tl.program_id(1)
q_start = tl.load(qo_indptr + batch_idx)
q_end = tl.load(qo_indptr + batch_idx + 1)
kv_start = tl.load(kv_indptr + batch_idx)
kv_end = tl.load(kv_indptr + batch_idx + 1)
seq_kv = kv_end - kv_start
actual_m = (q_end - q_start) * num_heads
if actual_m <= 0 or seq_kv <= 0: return
if m_block_idx * BLOCK_M >= actual_m: return
offs_m = m_block_idx * BLOCK_M + tl.arange(0, BLOCK_M)
mask_m = offs_m < actual_m
sq_idx = q_start + (offs_m // num_heads)
h_idx = offs_m % num_heads
q_ptrs_base = Q + sq_idx * stride_q_t + h_idx * stride_q_h
kv_scale_val = tl.load(kv_scale_ptr)
combined_scale = kv_scale_val * sm_scale
offs_d_lora = tl.arange(0, KV_LORA_RANK)
q_lora = tl.load(q_ptrs_base[:, None] + offs_d_lora[None, :] * stride_q_d, mask=mask_m[:, None], other=0.0)
offs_d_rope = KV_LORA_RANK + tl.arange(0, QK_ROPE_DIM)
q_rope = tl.load(q_ptrs_base[:, None] + offs_d_rope[None, :] * stride_q_d, mask=mask_m[:, None], other=0.0)
# -------------------------------------------------------------
# Quantize Q to FP8 on-the-fly to unlock MI355X FP8 Tensor Cores
# -------------------------------------------------------------
fp8_dtype = KV.dtype.element_ty
q_lora_amax = tl.maximum(tl.max(tl.abs(q_lora), axis=1), 1e-12)
q_lora_scale = q_lora_amax / 240.0
q_lora_fp8 = (q_lora / q_lora_scale[:, None]).to(fp8_dtype)
q_rope_amax = tl.maximum(tl.max(tl.abs(q_rope), axis=1), 1e-12)
q_rope_scale = q_rope_amax / 240.0
q_rope_fp8 = (q_rope / q_rope_scale[:, None]).to(fp8_dtype)
m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float('inf')
l_i = tl.zeros([BLOCK_M], dtype=tl.float32)
acc = tl.zeros([BLOCK_M, KV_LORA_RANK], dtype=tl.float32)
kv_base = KV + kv_start * stride_kv_t
offs_n = tl.arange(0, BLOCK_N)
kv_lora_ptrs = kv_base + offs_n[:, None] * stride_kv_t + offs_d_lora[None, :] * stride_kv_d
kv_rope_ptrs = kv_base + offs_n[:, None] * stride_kv_t + offs_d_rope[None, :] * stride_kv_d
for start_n in range(0, seq_kv, BLOCK_N):
curr_n = start_n + offs_n
mask_n = curr_n < seq_kv
kv_lora_fp8 = tl.load(kv_lora_ptrs, mask=mask_n[:, None], other=0.0)
kv_rope_fp8 = tl.load(kv_rope_ptrs, mask=mask_n[:, None], other=0.0)
# QK Matrix Multiply via FP8 Tensor Cores (4x faster than BF16)
qk = tl.dot(q_lora_fp8, tl.trans(kv_lora_fp8))
qk = qk.to(tl.float32) * q_lora_scale[:, None]
qk_rope = tl.dot(q_rope_fp8, tl.trans(kv_rope_fp8))
qk_rope = qk_rope.to(tl.float32) * q_rope_scale[:, None]
qk += qk_rope
qk = qk * combined_scale
qk = tl.where(mask_m[:, None] & mask_n[None, :], qk, float('-inf'))
m_ij = tl.max(qk, 1)
m_i_new = tl.maximum(m_i, m_ij)
alpha = tl.exp(m_i - m_i_new)
p = tl.exp(qk - m_i_new[:, None])
l_i_new = alpha * l_i + tl.sum(p, 1)
acc = acc * alpha[:, None]
kv_lora_bf16 = kv_lora_fp8.to(tl.bfloat16)
acc += tl.dot(p.to(tl.bfloat16), kv_lora_bf16)
m_i = m_i_new
l_i = l_i_new
kv_lora_ptrs += BLOCK_N * stride_kv_t
kv_rope_ptrs += BLOCK_N * stride_kv_t
acc = acc / l_i[:, None]
acc = acc * kv_scale_val
offs_d_v = tl.arange(0, KV_LORA_RANK)
mask_v = offs_d_v < v_head_dim
out_ptrs = Out + sq_idx[:, None]*stride_o_t + h_idx[:, None]*stride_o_h + offs_d_v[None, :]*stride_o_d
tl.store(out_ptrs, acc.to(Out.dtype.element_ty), mask=mask_m[:, None] & mask_v[None, :])
# ==============================================================================
# Split-KV Kernel
# ==============================================================================
@triton.autotune(
configs=[
# BLOCK_N=32: kv_lora tile (32,512) = 16+32=48 VGPRs vs (64,512) = 32+64=96 VGPRs
# Target: total VGPR ~120 → Occupancy=2 (2 blocks/CU)
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32}, num_stages=2, num_warps=4),
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32}, num_stages=1, num_warps=4),
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64}, num_stages=2, num_warps=4),
],
key=['KV_LORA_RANK', 'QK_ROPE_DIM', 'route_bucket'],
)
@triton.jit
def _mla_v11_split(
Q, KV, Partial_O, Partial_LSE,
qo_indptr, kv_indptr, kv_scale_ptr, sm_scale, route_bucket,
stride_q_t, stride_q_h, stride_q_d,
stride_kv_t, stride_kv_d,
stride_po_s, stride_po_t, stride_po_h, stride_po_d,
stride_plse_s, stride_plse_t, stride_plse_h,
num_heads: tl.constexpr, KV_LORA_RANK: tl.constexpr, QK_ROPE_DIM: tl.constexpr,
SPLIT_SIZE: tl.constexpr, NUM_KV_SPLITS: tl.constexpr,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
batch_idx = tl.program_id(0)
m_block_idx = tl.program_id(1)
split_idx = tl.program_id(2)
q_start = tl.load(qo_indptr + batch_idx)
q_end = tl.load(qo_indptr + batch_idx + 1)
kv_start = tl.load(kv_indptr + batch_idx)
kv_end = tl.load(kv_indptr + batch_idx + 1)
seq_kv = kv_end - kv_start
actual_m = (q_end - q_start) * num_heads
if actual_m <= 0 or seq_kv <= 0: return
if m_block_idx * BLOCK_M >= actual_m: return
offs_m = m_block_idx * BLOCK_M + tl.arange(0, BLOCK_M)
mask_m = offs_m < actual_m
sq_idx = q_start + (offs_m // num_heads)
h_idx = offs_m % num_heads
split_start = split_idx * SPLIT_SIZE
split_end = tl.minimum(split_start + SPLIT_SIZE, seq_kv)
split_len = split_end - split_start
if split_start >= seq_kv:
lse_ptrs = Partial_LSE + split_idx*stride_plse_s + sq_idx*stride_plse_t + h_idx*stride_plse_h
tl.store(lse_ptrs, float('-inf'), mask=mask_m)
return
q_ptrs_base = Q + sq_idx * stride_q_t + h_idx * stride_q_h
kv_scale_val = tl.load(kv_scale_ptr)
combined_scale = kv_scale_val * sm_scale
offs_d_lora = tl.arange(0, KV_LORA_RANK)
q_lora = tl.load(q_ptrs_base[:, None] + offs_d_lora[None, :] * stride_q_d, mask=mask_m[:, None], other=0.0)
offs_d_rope = KV_LORA_RANK + tl.arange(0, QK_ROPE_DIM)
q_rope = tl.load(q_ptrs_base[:, None] + offs_d_rope[None, :] * stride_q_d, mask=mask_m[:, None], other=0.0)
# -------------------------------------------------------------
# Quantize Q to FP8 on-the-fly to unlock MI355X FP8 Tensor Cores
# -------------------------------------------------------------
fp8_dtype = KV.dtype.element_ty
q_lora_amax = tl.maximum(tl.max(tl.abs(q_lora), axis=1), 1e-12)
q_lora_scale = q_lora_amax / 240.0
q_lora_fp8 = (q_lora / q_lora_scale[:, None]).to(fp8_dtype)
q_rope_amax = tl.maximum(tl.max(tl.abs(q_rope), axis=1), 1e-12)
q_rope_scale = q_rope_amax / 240.0
q_rope_fp8 = (q_rope / q_rope_scale[:, None]).to(fp8_dtype)
m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float('inf')
l_i = tl.zeros([BLOCK_M], dtype=tl.float32)
acc = tl.zeros([BLOCK_M, KV_LORA_RANK], dtype=tl.float32)
kv_base_ptr = KV + (kv_start + split_start) * stride_kv_t
offs_n = tl.arange(0, BLOCK_N)
kv_lora_ptrs = kv_base_ptr + offs_n[:, None] * stride_kv_t + offs_d_lora[None, :] * stride_kv_d
kv_rope_ptrs = kv_base_ptr + offs_n[:, None] * stride_kv_t + offs_d_rope[None, :] * stride_kv_d
for start_n in range(0, SPLIT_SIZE, BLOCK_N):
curr_n = start_n + offs_n
mask_n = curr_n < split_len
kv_lora_fp8 = tl.load(kv_lora_ptrs, mask=mask_n[:, None], other=0.0)
kv_rope_fp8 = tl.load(kv_rope_ptrs, mask=mask_n[:, None], other=0.0)
# QK Matrix Multiply via FP8 Tensor Cores
qk = tl.dot(q_lora_fp8, tl.trans(kv_lora_fp8))
qk = qk.to(tl.float32) * q_lora_scale[:, None]
qk_rope = tl.dot(q_rope_fp8, tl.trans(kv_rope_fp8))
qk_rope = qk_rope.to(tl.float32) * q_rope_scale[:, None]
qk += qk_rope
qk = qk * combined_scale
qk = tl.where(mask_m[:, None] & mask_n[None, :], qk, float('-inf'))
m_ij = tl.max(qk, 1)
m_i_new = tl.maximum(m_i, m_ij)
alpha = tl.exp(m_i - m_i_new)
p = tl.exp(qk - m_i_new[:, None])
l_i_new = alpha * l_i + tl.sum(p, 1)
acc = acc * alpha[:, None]
# Inline conversion: avoid named temporary to hint compiler for shorter liveness
acc += tl.dot(p.to(tl.bfloat16), kv_lora_fp8.to(tl.bfloat16))
m_i = m_i_new
l_i = l_i_new
kv_lora_ptrs += BLOCK_N * stride_kv_t
kv_rope_ptrs += BLOCK_N * stride_kv_t
acc = acc / l_i[:, None]
po_ptrs = Partial_O + split_idx*stride_po_s + sq_idx[:, None]*stride_po_t + h_idx[:, None]*stride_po_h + offs_d_lora[None, :]*stride_po_d
tl.store(po_ptrs, acc, mask=mask_m[:, None])
lse = m_i + tl.log(l_i)
lse_ptrs = Partial_LSE + split_idx*stride_plse_s + sq_idx*stride_plse_t + h_idx*stride_plse_h
tl.store(lse_ptrs, lse, mask=mask_m)
@triton.autotune(
configs=[
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64}, num_stages=1, num_warps=4),
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64}, num_stages=2, num_warps=4),
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 128}, num_stages=2, num_warps=4),
],
key=['KV_LORA_RANK', 'QK_ROPE_DIM', 'route_bucket'],
)
@triton.jit
def _mla_v28_short_no_split(
Q, KV, Out, qo_indptr, kv_indptr, kv_scale_ptr, sm_scale, route_bucket,
stride_q_t, stride_q_h, stride_q_d,
stride_kv_t, stride_kv_d,
stride_o_t, stride_o_h, stride_o_d,
num_heads: tl.constexpr, KV_LORA_RANK: tl.constexpr, QK_ROPE_DIM: tl.constexpr, v_head_dim: tl.constexpr,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
batch_idx = tl.program_id(0)
m_block_idx = tl.program_id(1)
q_start = tl.load(qo_indptr + batch_idx)
q_end = tl.load(qo_indptr + batch_idx + 1)
kv_start = tl.load(kv_indptr + batch_idx)
kv_end = tl.load(kv_indptr + batch_idx + 1)
seq_kv = kv_end - kv_start
actual_m = (q_end - q_start) * num_heads
if actual_m <= 0 or seq_kv <= 0: return
if m_block_idx * BLOCK_M >= actual_m: return
offs_m = m_block_idx * BLOCK_M + tl.arange(0, BLOCK_M)
mask_m = offs_m < actual_m
sq_idx = q_start + (offs_m // num_heads)
h_idx = offs_m % num_heads
q_ptrs_base = Q + sq_idx * stride_q_t + h_idx * stride_q_h
kv_scale_val = tl.load(kv_scale_ptr)
combined_scale = kv_scale_val * sm_scale
offs_d_lora = tl.arange(0, KV_LORA_RANK)
q_lora = tl.load(q_ptrs_base[:, None] + offs_d_lora[None, :] * stride_q_d, mask=mask_m[:, None], other=0.0)
offs_d_rope = KV_LORA_RANK + tl.arange(0, QK_ROPE_DIM)
q_rope = tl.load(q_ptrs_base[:, None] + offs_d_rope[None, :] * stride_q_d, mask=mask_m[:, None], other=0.0)
fp8_dtype = KV.dtype.element_ty
q_lora_amax = tl.maximum(tl.max(tl.abs(q_lora), axis=1), 1e-12)
q_lora_scale = q_lora_amax / 240.0
q_lora_fp8 = (q_lora / q_lora_scale[:, None]).to(fp8_dtype)
q_rope_amax = tl.maximum(tl.max(tl.abs(q_rope), axis=1), 1e-12)
q_rope_scale = q_rope_amax / 240.0
q_rope_fp8 = (q_rope / q_rope_scale[:, None]).to(fp8_dtype)
m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float('inf')
l_i = tl.zeros([BLOCK_M], dtype=tl.float32)
acc = tl.zeros([BLOCK_M, KV_LORA_RANK], dtype=tl.float32)
kv_base = KV + kv_start * stride_kv_t
offs_n = tl.arange(0, BLOCK_N)
kv_lora_ptrs = kv_base + offs_n[:, None] * stride_kv_t + offs_d_lora[None, :] * stride_kv_d
kv_rope_ptrs = kv_base + offs_n[:, None] * stride_kv_t + offs_d_rope[None, :] * stride_kv_d
for start_n in range(0, seq_kv, BLOCK_N):
curr_n = start_n + offs_n
mask_n = curr_n < seq_kv
kv_lora_fp8 = tl.load(kv_lora_ptrs, mask=mask_n[:, None], other=0.0)
kv_rope_fp8 = tl.load(kv_rope_ptrs, mask=mask_n[:, None], other=0.0)
qk = tl.dot(q_lora_fp8, tl.trans(kv_lora_fp8))
qk = qk.to(tl.float32) * q_lora_scale[:, None]
qk_rope = tl.dot(q_rope_fp8, tl.trans(kv_rope_fp8))
qk_rope = qk_rope.to(tl.float32) * q_rope_scale[:, None]
qk += qk_rope
qk = qk * combined_scale
qk = tl.where(mask_m[:, None] & mask_n[None, :], qk, float('-inf'))
m_ij = tl.max(qk, 1)
m_i_new = tl.maximum(m_i, m_ij)
alpha = tl.exp(m_i - m_i_new)
p = tl.exp(qk - m_i_new[:, None])
l_i_new = alpha * l_i + tl.sum(p, 1)
acc = acc * alpha[:, None]
acc += tl.dot(p.to(tl.bfloat16), kv_lora_fp8.to(tl.bfloat16))
m_i = m_i_new
l_i = l_i_new
kv_lora_ptrs += BLOCK_N * stride_kv_t
kv_rope_ptrs += BLOCK_N * stride_kv_t
acc = acc / l_i[:, None]
acc = acc * kv_scale_val
offs_d_v = tl.arange(0, KV_LORA_RANK)
mask_v = offs_d_v < v_head_dim
out_ptrs = Out + sq_idx[:, None]*stride_o_t + h_idx[:, None]*stride_o_h + offs_d_v[None, :]*stride_o_d
tl.store(out_ptrs, acc.to(Out.dtype.element_ty), mask=mask_m[:, None] & mask_v[None, :])
@triton.autotune(
configs=[
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32}, num_stages=1, num_warps=4),
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32}, num_stages=2, num_warps=4),
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64}, num_stages=1, num_warps=4),
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64}, num_stages=2, num_warps=4),
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 128}, num_stages=2, num_warps=4),
],
key=['KV_LORA_RANK', 'QK_ROPE_DIM', 'route_bucket'],
)
@triton.jit
def _mla_v28_short_split(
Q, KV, Partial_O, Partial_LSE,
qo_indptr, kv_indptr, kv_scale_ptr, sm_scale, route_bucket,
stride_q_t, stride_q_h, stride_q_d,
stride_kv_t, stride_kv_d,
stride_po_s, stride_po_t, stride_po_h, stride_po_d,
stride_plse_s, stride_plse_t, stride_plse_h,
num_heads: tl.constexpr, KV_LORA_RANK: tl.constexpr, QK_ROPE_DIM: tl.constexpr,
SPLIT_SIZE: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
batch_idx = tl.program_id(0)
m_block_idx = tl.program_id(1)
split_idx = tl.program_id(2)
q_start = tl.load(qo_indptr + batch_idx)
q_end = tl.load(qo_indptr + batch_idx + 1)
kv_start = tl.load(kv_indptr + batch_idx)
kv_end = tl.load(kv_indptr + batch_idx + 1)
seq_kv = kv_end - kv_start
actual_m = (q_end - q_start) * num_heads
if actual_m <= 0 or seq_kv <= 0: return
if m_block_idx * BLOCK_M >= actual_m: return
offs_m = m_block_idx * BLOCK_M + tl.arange(0, BLOCK_M)
mask_m = offs_m < actual_m
sq_idx = q_start + (offs_m // num_heads)
h_idx = offs_m % num_heads
split_start = split_idx * SPLIT_SIZE
split_end = tl.minimum(split_start + SPLIT_SIZE, seq_kv)
split_len = split_end - split_start
if split_start >= seq_kv:
lse_ptrs = Partial_LSE + split_idx*stride_plse_s + sq_idx*stride_plse_t + h_idx*stride_plse_h
tl.store(lse_ptrs, float('-inf'), mask=mask_m)
return
q_ptrs_base = Q + sq_idx * stride_q_t + h_idx * stride_q_h
kv_scale_val = tl.load(kv_scale_ptr)
combined_scale = kv_scale_val * sm_scale
offs_d_lora = tl.arange(0, KV_LORA_RANK)
q_lora = tl.load(q_ptrs_base[:, None] + offs_d_lora[None, :] * stride_q_d, mask=mask_m[:, None], other=0.0)
offs_d_rope = KV_LORA_RANK + tl.arange(0, QK_ROPE_DIM)
q_rope = tl.load(q_ptrs_base[:, None] + offs_d_rope[None, :] * stride_q_d, mask=mask_m[:, None], other=0.0)
fp8_dtype = KV.dtype.element_ty
q_lora_amax = tl.maximum(tl.max(tl.abs(q_lora), axis=1), 1e-12)
q_lora_scale = q_lora_amax / 240.0
q_lora_fp8 = (q_lora / q_lora_scale[:, None]).to(fp8_dtype)
q_rope_amax = tl.maximum(tl.max(tl.abs(q_rope), axis=1), 1e-12)
q_rope_scale = q_rope_amax / 240.0
q_rope_fp8 = (q_rope / q_rope_scale[:, None]).to(fp8_dtype)
m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float('inf')
l_i = tl.zeros([BLOCK_M], dtype=tl.float32)
acc = tl.zeros([BLOCK_M, KV_LORA_RANK], dtype=tl.float32)
kv_base_ptr = KV + (kv_start + split_start) * stride_kv_t
offs_n = tl.arange(0, BLOCK_N)
kv_lora_ptrs = kv_base_ptr + offs_n[:, None] * stride_kv_t + offs_d_lora[None, :] * stride_kv_d
kv_rope_ptrs = kv_base_ptr + offs_n[:, None] * stride_kv_t + offs_d_rope[None, :] * stride_kv_d
for start_n in range(0, SPLIT_SIZE, BLOCK_N):
curr_n = start_n + offs_n
mask_n = curr_n < split_len
kv_lora_fp8 = tl.load(kv_lora_ptrs, mask=mask_n[:, None], other=0.0)
kv_rope_fp8 = tl.load(kv_rope_ptrs, mask=mask_n[:, None], other=0.0)
qk = tl.dot(q_lora_fp8, tl.trans(kv_lora_fp8))
qk = qk.to(tl.float32) * q_lora_scale[:, None]
qk_rope = tl.dot(q_rope_fp8, tl.trans(kv_rope_fp8))
qk_rope = qk_rope.to(tl.float32) * q_rope_scale[:, None]
qk += qk_rope
qk = qk * combined_scale
qk = tl.where(mask_m[:, None] & mask_n[None, :], qk, float('-inf'))
m_ij = tl.max(qk, 1)
m_i_new = tl.maximum(m_i, m_ij)
alpha = tl.exp(m_i - m_i_new)
p = tl.exp(qk - m_i_new[:, None])
l_i_new = alpha * l_i + tl.sum(p, 1)
acc = acc * alpha[:, None]
acc += tl.dot(p.to(tl.bfloat16), kv_lora_fp8.to(tl.bfloat16))
m_i = m_i_new
l_i = l_i_new
kv_lora_ptrs += BLOCK_N * stride_kv_t
kv_rope_ptrs += BLOCK_N * stride_kv_t
acc = acc / l_i[:, None]
po_ptrs = Partial_O + split_idx*stride_po_s + sq_idx[:, None]*stride_po_t + h_idx[:, None]*stride_po_h + offs_d_lora[None, :]*stride_po_d
tl.store(po_ptrs, acc, mask=mask_m[:, None])
lse = m_i + tl.log(l_i)
lse_ptrs = Partial_LSE + split_idx*stride_plse_s + sq_idx*stride_plse_t + h_idx*stride_plse_h
tl.store(lse_ptrs, lse, mask=mask_m)
@triton.jit
def _reduce_kernel(
Partial_O, Partial_LSE, Out, kv_scale_ptr, total_th,
stride_po_s, stride_po_t, stride_po_h, stride_po_d,
stride_plse_s, stride_plse_t, stride_plse_h,
stride_o_t, stride_o_h, stride_o_d,
NUM_HEADS: tl.constexpr, NUM_KV_SPLITS: tl.constexpr, V_HEAD_DIM: tl.constexpr, BLOCK_D: tl.constexpr,
):
pid = tl.program_id(0)
if pid >= total_th: return
t_idx = pid // NUM_HEADS
h_idx = pid % NUM_HEADS
kv_scale_val = tl.load(kv_scale_ptr)
offs_s = tl.arange(0, NUM_KV_SPLITS)
lse_ptrs = Partial_LSE + offs_s*stride_plse_s + t_idx*stride_plse_t + h_idx*stride_plse_h
lse_vals = tl.load(lse_ptrs)
lse_max = tl.max(lse_vals)
lse_max = tl.where(lse_max == float('-inf'), 0.0, lse_max)
weights = tl.exp(lse_vals - lse_max)
w_sum = tl.sum(weights)
for d_start in range(0, V_HEAD_DIM, BLOCK_D):
offs_d = d_start + tl.arange(0, BLOCK_D)
mask_d = offs_d < V_HEAD_DIM
po_ptrs = Partial_O + offs_s[:, None]*stride_po_s + t_idx*stride_po_t + h_idx*stride_po_h + offs_d[None, :]*stride_po_d
po_vals = tl.load(po_ptrs, mask=mask_d[None, :], other=0.0)
weighted = weights[:, None] * po_vals
acc = tl.sum(weighted, axis=0)
acc = acc / tl.maximum(w_sum, 1e-12) * kv_scale_val
out_ptrs = Out + t_idx*stride_o_t + h_idx*stride_o_h + offs_d*stride_o_d
tl.store(out_ptrs, acc.to(tl.bfloat16), mask=mask_d)
@triton.jit
def _reduce_kernel_dsplit(
Partial_O, Partial_LSE, Out, kv_scale_ptr, total_th,
stride_po_s, stride_po_t, stride_po_h, stride_po_d,
stride_plse_s, stride_plse_t, stride_plse_h,
stride_o_t, stride_o_h, stride_o_d,
NUM_HEADS: tl.constexpr, NUM_KV_SPLITS: tl.constexpr, V_HEAD_DIM: tl.constexpr, BLOCK_D: tl.constexpr,
):
pid_th = tl.program_id(0)
pid_d = tl.program_id(1)
if pid_th >= total_th:
return
t_idx = pid_th // NUM_HEADS
h_idx = pid_th % NUM_HEADS
kv_scale_val = tl.load(kv_scale_ptr)
offs_s = tl.arange(0, NUM_KV_SPLITS)
lse_ptrs = Partial_LSE + offs_s*stride_plse_s + t_idx*stride_plse_t + h_idx*stride_plse_h
lse_vals = tl.load(lse_ptrs)
lse_max = tl.max(lse_vals)
lse_max = tl.where(lse_max == float('-inf'), 0.0, lse_max)
weights = tl.exp(lse_vals - lse_max)
scale = kv_scale_val / tl.maximum(tl.sum(weights), 1e-12)
d_start = pid_d * BLOCK_D
offs_d = d_start + tl.arange(0, BLOCK_D)
mask_d = offs_d < V_HEAD_DIM
po_ptrs = Partial_O + offs_s[:, None]*stride_po_s + t_idx*stride_po_t + h_idx*stride_po_h + offs_d[None, :]*stride_po_d
po_vals = tl.load(po_ptrs, mask=mask_d[None, :], other=0.0)
acc = tl.sum(weights[:, None] * po_vals, axis=0)
acc = acc * scale
out_ptrs = Out + t_idx*stride_o_t + h_idx*stride_o_h + offs_d*stride_o_d
tl.store(out_ptrs, acc.to(tl.bfloat16), mask=mask_d)
def _short_kv_target_blocks(batch_size: int) -> int:
if batch_size <= 4:
return SHORT_KV_TARGET_BLOCKS_SMALL
return SHORT_KV_TARGET_BLOCKS
def _short_split_route_bucket(split_size: int) -> int:
if split_size <= 64:
return 0
if split_size <= 128:
return 1
return 2
def _long_split_route_bucket(num_kv_splits: int) -> int:
if num_kv_splits >= 32:
return 0
if num_kv_splits >= 8:
return 1
return 2
# ==============================================================================
# v29: 短 KV 独立 autotune cache + 长 KV split 分桶
# ==============================================================================
def custom_kernel_v29(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
num_heads = config["num_heads"]
qk_head_dim = config["qk_head_dim"]
v_head_dim = config["v_head_dim"]
sm_scale = config["sm_scale"]
kv_lora_rank = config.get("kv_lora_rank", 512)
qk_rope_head_dim = config.get("qk_rope_head_dim", 64)
kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
kv_fp8_2d = kv_buffer_fp8.view(-1, qk_head_dim)
batch_size = qo_indptr.shape[0] - 1
total_q = q.shape[0]
total_kv = kv_buffer_fp8.shape[0]
q_seq_len = int(config.get("q_seq_len", max(1, (total_q + batch_size - 1) // batch_size)))
kv_seq_len = int(config.get("kv_seq_len", max(1, (total_kv + batch_size - 1) // batch_size)))
short_kv_special = (
q_seq_len == 1
and kv_seq_len <= SHORT_KV_MAX
)
if not short_kv_special:
return custom_kernel_v14(data)
max_m = q_seq_len * num_heads
block_m_est = 16
m_blocks = (max_m + block_m_est - 1) // block_m_est
base_blocks = batch_size * m_blocks
target_blocks = _short_kv_target_blocks(batch_size)
use_split_kv = base_blocks < target_blocks
if not use_split_kv:
out = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device=q.device)
def grid_fn(META): return (batch_size, triton.cdiv(max_m, META['BLOCK_M']))
_mla_v28_short_no_split[grid_fn](
q, kv_fp8_2d, out, qo_indptr, kv_indptr, kv_scale_fp8, sm_scale, 0,
q.stride(0), q.stride(1), q.stride(2),
kv_fp8_2d.stride(0), kv_fp8_2d.stride(1),
out.stride(0), out.stride(1), out.stride(2),
num_heads=num_heads, KV_LORA_RANK=kv_lora_rank,
QK_ROPE_DIM=qk_rope_head_dim, v_head_dim=v_head_dim,
)
return out
needed_splits = max(1, (target_blocks + base_blocks - 1) // base_blocks)
num_kv_splits = 1
while num_kv_splits < needed_splits:
num_kv_splits *= 2
num_kv_splits = min(num_kv_splits, SHORT_KV_MAX_SPLITS)
raw_split = (kv_seq_len + num_kv_splits - 1) // num_kv_splits
BLOCK_N_max = 32
split_size = max(BLOCK_N_max, ((raw_split + BLOCK_N_max - 1) // BLOCK_N_max) * BLOCK_N_max)
route_bucket = _short_split_route_bucket(split_size)
partial_o = torch.empty(
(num_kv_splits, total_q, num_heads, kv_lora_rank),
dtype=torch.float32, device=q.device,
)
partial_lse = torch.empty(
(num_kv_splits, total_q, num_heads),
dtype=torch.float32, device=q.device,
)
def grid_fn_split(META):
return (batch_size, triton.cdiv(max_m, META['BLOCK_M']), num_kv_splits)
_mla_v28_short_split[grid_fn_split](
q, kv_fp8_2d, partial_o, partial_lse,
qo_indptr, kv_indptr, kv_scale_fp8, sm_scale, route_bucket,
q.stride(0), q.stride(1), q.stride(2),
kv_fp8_2d.stride(0), kv_fp8_2d.stride(1),
partial_o.stride(0), partial_o.stride(1), partial_o.stride(2), partial_o.stride(3),
partial_lse.stride(0), partial_lse.stride(1), partial_lse.stride(2),
num_heads=num_heads, KV_LORA_RANK=kv_lora_rank, QK_ROPE_DIM=qk_rope_head_dim,
SPLIT_SIZE=split_size,
)
out = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device=q.device)
total_th = total_q * num_heads
_reduce_kernel[(total_th,)](
partial_o, partial_lse, out, kv_scale_fp8, total_th,
partial_o.stride(0), partial_o.stride(1), partial_o.stride(2), partial_o.stride(3),
partial_lse.stride(0), partial_lse.stride(1), partial_lse.stride(2),
out.stride(0), out.stride(1), out.stride(2),
NUM_HEADS=num_heads, NUM_KV_SPLITS=num_kv_splits, V_HEAD_DIM=v_head_dim,
BLOCK_D=128, num_warps=1,
)
return out
# ==============================================================================
# v41: 在 v37 基础上只测试 small case 的 partial_o 带宽压缩
# ==============================================================================
def custom_kernel_v41(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
num_heads = config["num_heads"]
qk_head_dim = config["qk_head_dim"]
v_head_dim = config["v_head_dim"]
sm_scale = config["sm_scale"]
kv_lora_rank = config.get("kv_lora_rank", 512)
qk_rope_head_dim = config.get("qk_rope_head_dim", 64)
kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
kv_fp8_2d = kv_buffer_fp8.view(-1, qk_head_dim)
batch_size = qo_indptr.shape[0] - 1
total_q = q.shape[0]
total_kv = kv_buffer_fp8.shape[0]
q_seq_len = int(config.get("q_seq_len", max(1, (total_q + batch_size - 1) // batch_size)))
kv_seq_len = int(config.get("kv_seq_len", max(1, (total_kv + batch_size - 1) // batch_size)))
short_kv_special = (
q_seq_len == 1
and kv_seq_len <= SHORT_KV_MAX
)
if not short_kv_special:
return custom_kernel_v14(data)
max_m = q_seq_len * num_heads
block_m_est = 16
m_blocks = (max_m + block_m_est - 1) // block_m_est
base_blocks = batch_size * m_blocks
target_blocks = _short_kv_target_blocks(batch_size)
use_split_kv = base_blocks < target_blocks
if not use_split_kv:
out = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device=q.device)
def grid_fn(META): return (batch_size, triton.cdiv(max_m, META['BLOCK_M']))
_mla_v28_short_no_split[grid_fn](
q, kv_fp8_2d, out, qo_indptr, kv_indptr, kv_scale_fp8, sm_scale, 0,
q.stride(0), q.stride(1), q.stride(2),
kv_fp8_2d.stride(0), kv_fp8_2d.stride(1),
out.stride(0), out.stride(1), out.stride(2),
num_heads=num_heads, KV_LORA_RANK=kv_lora_rank,
QK_ROPE_DIM=qk_rope_head_dim, v_head_dim=v_head_dim,
)
return out
needed_splits = max(1, (target_blocks + base_blocks - 1) // base_blocks)
num_kv_splits = 1
while num_kv_splits < needed_splits:
num_kv_splits *= 2
num_kv_splits = min(num_kv_splits, SHORT_KV_MAX_SPLITS)
raw_split = (kv_seq_len + num_kv_splits - 1) // num_kv_splits
BLOCK_N_max = 32
split_size = max(BLOCK_N_max, ((raw_split + BLOCK_N_max - 1) // BLOCK_N_max) * BLOCK_N_max)
route_bucket = _short_split_route_bucket(split_size)
total_th = total_q * num_heads
use_small_bf16_partial_o = total_th <= SHORT_KV_SMALL_REDUCE_MAX_TOTAL_TH
partial_o_dtype = torch.bfloat16 if use_small_bf16_partial_o else torch.float32
partial_o = torch.empty(
(num_kv_splits, total_q, num_heads, kv_lora_rank),
dtype=partial_o_dtype, device=q.device,
)
partial_lse = torch.empty(
(num_kv_splits, total_q, num_heads),
dtype=torch.float32, device=q.device,
)
def grid_fn_split(META):
return (batch_size, triton.cdiv(max_m, META['BLOCK_M']), num_kv_splits)
_mla_v28_short_split[grid_fn_split](
q, kv_fp8_2d, partial_o, partial_lse,
qo_indptr, kv_indptr, kv_scale_fp8, sm_scale, route_bucket,
q.stride(0), q.stride(1), q.stride(2),
kv_fp8_2d.stride(0), kv_fp8_2d.stride(1),
partial_o.stride(0), partial_o.stride(1), partial_o.stride(2), partial_o.stride(3),
partial_lse.stride(0), partial_lse.stride(1), partial_lse.stride(2),
num_heads=num_heads, KV_LORA_RANK=kv_lora_rank, QK_ROPE_DIM=qk_rope_head_dim,
SPLIT_SIZE=split_size,
)
out = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device=q.device)
use_small_dsplit_reduce = (
total_th <= SHORT_KV_SMALL_REDUCE_MAX_TOTAL_TH
and v_head_dim >= SHORT_KV_SMALL_REDUCE_BLOCK_D
)
if use_small_dsplit_reduce:
def grid_fn_reduce(META):
return (total_th, triton.cdiv(v_head_dim, META['BLOCK_D']))
_reduce_kernel_dsplit[grid_fn_reduce](
partial_o, partial_lse, out, kv_scale_fp8, total_th,
partial_o.stride(0), partial_o.stride(1), partial_o.stride(2), partial_o.stride(3),
partial_lse.stride(0), partial_lse.stride(1), partial_lse.stride(2),
out.stride(0), out.stride(1), out.stride(2),
NUM_HEADS=num_heads, NUM_KV_SPLITS=num_kv_splits, V_HEAD_DIM=v_head_dim,
BLOCK_D=SHORT_KV_SMALL_REDUCE_BLOCK_D, num_warps=1,
)
return out
_reduce_kernel[(total_th,)](
partial_o, partial_lse, out, kv_scale_fp8, total_th,
partial_o.stride(0), partial_o.stride(1), partial_o.stride(2), partial_o.stride(3),
partial_lse.stride(0), partial_lse.stride(1), partial_lse.stride(2),
out.stride(0), out.stride(1), out.stride(2),
NUM_HEADS=num_heads, NUM_KV_SPLITS=num_kv_splits, V_HEAD_DIM=v_head_dim,
BLOCK_D=128, num_warps=1,
)
return out
# ==============================================================================
# v14: BLOCK_N=32 压缩 VGPR → Occupancy=2
# ==============================================================================
def custom_kernel_v14(data: input_t) -> output_t:
q, kv_data, qo_indptr, kv_indptr, config = data
num_heads = config["num_heads"]
qk_head_dim = config["qk_head_dim"] # 576
v_head_dim = config["v_head_dim"] # 512
sm_scale = config["sm_scale"]
kv_lora_rank = config.get("kv_lora_rank", 512)
qk_rope_head_dim = config.get("qk_rope_head_dim", 64)
kv_buffer_fp8, kv_scale_fp8 = kv_data["fp8"]
kv_fp8_2d = kv_buffer_fp8.view(-1, qk_head_dim)
batch_size = qo_indptr.shape[0] - 1
total_q = q.shape[0]
total_kv = kv_buffer_fp8.shape[0]
q_seq_avg = (total_q + batch_size - 1) // batch_size
q_seq_len_est = 4 if q_seq_avg <= 4 else q_seq_avg + 4 # Padding bound
max_m = q_seq_len_est * num_heads
kv_seq_avg = (total_kv + batch_size - 1) // batch_size
kv_seq_max = kv_seq_avg + 384
# ===================================
# V11 激进优化:取消短序列 No-Split 惩罚,允许极端切片
# ===================================
BLOCK_M_est = 16
m_blocks = (max_m + BLOCK_M_est - 1) // BLOCK_M_est
base_blocks = batch_size * m_blocks
# 只要 base_blocks < 1024 (为了撑满极高 VGPR 所需的 CU Occupancy),就必须切分!
use_split_kv = base_blocks < 1024
if not use_split_kv:
out = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device=q.device)
def grid_fn(META): return (batch_size, triton.cdiv(max_m, META['BLOCK_M']))
_mla_v11_no_split[grid_fn](
q, kv_fp8_2d, out, qo_indptr, kv_indptr, kv_scale_fp8, sm_scale,
q.stride(0), q.stride(1), q.stride(2),
kv_fp8_2d.stride(0), kv_fp8_2d.stride(1),
out.stride(0), out.stride(1), out.stride(2),
num_heads=num_heads, KV_LORA_RANK=kv_lora_rank,
QK_ROPE_DIM=qk_rope_head_dim, v_head_dim=v_head_dim,
)
return out
# ====== Split-KV =======
# 由于我们需要保障每一个 CU 尽可能运转,目标总块数达到 1024
needed_splits = max(1, (1024 + base_blocks - 1) // base_blocks)
num_kv_splits = 1
while num_kv_splits < needed_splits:
num_kv_splits *= 2
# 彻底解除切分粒度封印,但最高至 32 以防 Reduce 发生降速反噬
num_kv_splits = min(num_kv_splits, 32)
raw_split = (kv_seq_max + num_kv_splits - 1) // num_kv_splits
BLOCK_N_max = 32
split_size = max(BLOCK_N_max, ((raw_split + BLOCK_N_max - 1) // BLOCK_N_max) * BLOCK_N_max)
partial_o = torch.empty(
(num_kv_splits, total_q, num_heads, kv_lora_rank),
dtype=torch.float32, device=q.device,
)
partial_lse = torch.empty(
(num_kv_splits, total_q, num_heads),
dtype=torch.float32, device=q.device,
)
def grid_fn_split(META):
return (batch_size, triton.cdiv(max_m, META['BLOCK_M']), num_kv_splits)
long_route_bucket = _long_split_route_bucket(num_kv_splits)
_mla_v11_split[grid_fn_split](
q, kv_fp8_2d, partial_o, partial_lse,
qo_indptr, kv_indptr, kv_scale_fp8, sm_scale, long_route_bucket,
q.stride(0), q.stride(1), q.stride(2),
kv_fp8_2d.stride(0), kv_fp8_2d.stride(1),
partial_o.stride(0), partial_o.stride(1), partial_o.stride(2), partial_o.stride(3),
partial_lse.stride(0), partial_lse.stride(1), partial_lse.stride(2),
num_heads=num_heads, KV_LORA_RANK=kv_lora_rank, QK_ROPE_DIM=qk_rope_head_dim,
SPLIT_SIZE=split_size, NUM_KV_SPLITS=num_kv_splits,
)
out = torch.empty((total_q, num_heads, v_head_dim), dtype=torch.bfloat16, device=q.device)
total_th = total_q * num_heads
_reduce_kernel[(total_th,)](
partial_o, partial_lse, out, kv_scale_fp8, total_th,
partial_o.stride(0), partial_o.stride(1), partial_o.stride(2), partial_o.stride(3),
partial_lse.stride(0), partial_lse.stride(1), partial_lse.stride(2),
out.stride(0), out.stride(1), out.stride(2),
NUM_HEADS=num_heads, NUM_KV_SPLITS=num_kv_splits, V_HEAD_DIM=v_head_dim,
BLOCK_D=128, num_warps=1,
)
return out
scrolls · 939 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