gpt-5 / tritona41cd4
gpt-5_triton_a41cd4 · gpt-5-2025-08-07 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 286 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-a41cd4?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32
Benchmark evidence
47 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=8
NVIDIA B200
101.2µs
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=108
NVIDIA B200
358.4µs
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=457
NVIDIA B200
374.7µs
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=208
NVIDIA B200
617.3µs
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=308
NVIDIA B200
870.8µs
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=408
NVIDIA B200
1.12ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=1857
NVIDIA B200
1.25ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=508
NVIDIA B200
1.38ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=7257
NVIDIA B200
1.48ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=608
NVIDIA B200
1.64ms
#4 of 4
2025-10-16
Show all 47 measurements ›Showing all 47 measurements ⌄
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=8057
NVIDIA B200
1.76ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=5057
NVIDIA B200
1.85ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=708
NVIDIA B200
1.92ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=5857
NVIDIA B200
1.97ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=6657
NVIDIA B200
2.13ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=808
NVIDIA B200
2.18ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=1008
NVIDIA B200
2.70ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=1108
NVIDIA B200
2.95ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=1208
NVIDIA B200
3.21ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=4357
NVIDIA B200
3.55ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=2757
NVIDIA B200
3.72ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=3557
NVIDIA B200
3.86ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=8857
NVIDIA B200
4.24ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=9945
NVIDIA B200
4.57ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=16345
NVIDIA B200
4.83ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=1908
NVIDIA B200
5.04ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=22745
NVIDIA B200
5.04ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=27545
NVIDIA B200
5.27ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=9657
NVIDIA B200
5.41ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=30745
NVIDIA B200
5.46ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=33945
NVIDIA B200
5.52ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=10857
NVIDIA B200
5.62ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=37145
NVIDIA B200
5.72ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=40345
NVIDIA B200
5.79ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=12857
NVIDIA B200
6.02ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=2408
NVIDIA B200
6.34ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=14857
NVIDIA B200
6.39ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [16, 16, 64] · num_kv_indices=17257
NVIDIA B200
6.95ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [1, 16, 64] · num_kv_indices=2708
NVIDIA B200
7.11ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=44845
NVIDIA B200
8.98ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=48045
NVIDIA B200
9.15ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=51245
NVIDIA B200
9.30ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=54445
NVIDIA B200
9.44ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=57645
NVIDIA B200
9.61ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=62345
NVIDIA B200
23.5ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=75145
NVIDIA B200
24.0ms
#4 of 4
2025-10-16
MLA paged decode h16 ckv512 kpe64 ps1bf16 · [64, 16, 64] · num_kv_indices=68745
NVIDIA B200
27.3ms
#4 of 4
2025-10-16
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:c4d597804add5bfbd3c24de72e5a7c66c9f333e568fe8cdfcc9b7346fcc4cbd0
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-5-2025-08-07
imported2026-08-20
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
main.py286 lines
import math
import torch
import triton
import triton.language as tl
@triton.jit
def mla_paged_decode_h16_ckv512_kpe64_ps1_kernel(
q_nope_ptr,
q_pe_ptr,
ckv_ptr,
kpe_ptr,
kv_indptr_ptr,
kv_indices_ptr,
output_ptr,
lse_ptr,
B,
H: tl.constexpr,
DCKV: tl.constexpr,
DKPE: tl.constexpr,
sm_scale,
stride_qn_b,
stride_qn_h,
stride_qn_d,
stride_qp_b,
stride_qp_h,
stride_qp_d,
stride_ckv_p,
stride_ckv_d,
stride_kpe_p,
stride_kpe_d,
stride_out_b,
stride_out_h,
stride_out_d,
stride_lse_b,
stride_lse_h,
BLOCK_TOK: tl.constexpr,
BLOCK_DCKV: tl.constexpr,
BLOCK_DKPE: tl.constexpr,
):
pid = tl.program_id(0)
b = pid // H
h = pid % H
if b >= B:
return
page_start = tl.load(kv_indptr_ptr + b, mask=True, other=0).to(tl.int32)
page_end = tl.load(kv_indptr_ptr + (b + 1), mask=True, other=0).to(tl.int32)
L = page_end - page_start
qn_base = q_nope_ptr + b * stride_qn_b + h * stride_qn_h
qp_base = q_pe_ptr + b * stride_qp_b + h * stride_qp_h
out_base = output_ptr + b * stride_out_b + h * stride_out_h
lse_off = lse_ptr + b * stride_lse_b + h * stride_lse_h
# Early exit
if L <= 0:
offs_d = tl.arange(0, BLOCK_DCKV)
zero_bf16 = tl.zeros([BLOCK_DCKV], dtype=tl.bfloat16)
for t in range(0, DCKV, BLOCK_DCKV):
d = t + offs_d
mask_d = d < DCKV
tl.store(out_base + d * stride_out_d, zero_bf16, mask=mask_d)
tl.store(lse_off, -float("inf"))
return
# Preload q vectors
offs_dckv = tl.arange(0, BLOCK_DCKV)
qn0 = tl.load(qn_base + offs_dckv * stride_qn_d, mask=offs_dckv < DCKV, other=0.0).to(tl.float32)
d1 = offs_dckv + BLOCK_DCKV
qn1 = tl.load(qn_base + d1 * stride_qn_d, mask=d1 < DCKV, other=0.0).to(tl.float32)
d2 = offs_dckv + 2 * BLOCK_DCKV
qn2 = tl.load(qn_base + d2 * stride_qn_d, mask=d2 < DCKV, other=0.0).to(tl.float32)
d3 = offs_dckv + 3 * BLOCK_DCKV
qn3 = tl.load(qn_base + d3 * stride_qn_d, mask=d3 < DCKV, other=0.0).to(tl.float32)
offs_kpe = tl.arange(0, BLOCK_DKPE)
qp_vec = tl.load(qp_base + offs_kpe * stride_qp_d, mask=offs_kpe < DKPE, other=0.0).to(tl.float32)
# Streaming softmax stats (natural log domain)
m = tl.full([], -float("inf"), dtype=tl.float32)
S = tl.full([], 0.0, dtype=tl.float32)
# Numerator accumulators for output (float32)
O0 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
O1 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
O2 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
O3 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
start = tl.zeros([], dtype=tl.int32)
while start < L:
# Process BLOCK_TOK tokens sequentially for better numerical stability
for i in range(BLOCK_TOK):
t = start + i
valid = t < L
# Load token index
tok = tl.load(kv_indices_ptr + page_start + t, mask=valid, other=0).to(tl.int32)
# Compute logits for this token
# Kc dot qn
K0 = tl.load(
ckv_ptr + tok * stride_ckv_p + offs_dckv * stride_ckv_d,
mask=valid & (offs_dckv < DCKV),
other=0.0,
).to(tl.float32)
l = tl.sum(K0 * qn0, axis=0)
K1 = tl.load(
ckv_ptr + tok * stride_ckv_p + d1 * stride_ckv_d,
mask=valid & (d1 < DCKV),
other=0.0,
).to(tl.float32)
l += tl.sum(K1 * qn1, axis=0)
K2 = tl.load(
ckv_ptr + tok * stride_ckv_p + d2 * stride_ckv_d,
mask=valid & (d2 < DCKV),
other=0.0,
).to(tl.float32)
l += tl.sum(K2 * qn2, axis=0)
K3 = tl.load(
ckv_ptr + tok * stride_ckv_p + d3 * stride_ckv_d,
mask=valid & (d3 < DCKV),
other=0.0,
).to(tl.float32)
l += tl.sum(K3 * qn3, axis=0)
# Kp dot qp
KP = tl.load(
kpe_ptr + tok * stride_kpe_p + offs_kpe * stride_kpe_d,
mask=valid & (offs_kpe < DKPE),
other=0.0,
).to(tl.float32)
l += tl.sum(KP * qp_vec, axis=0)
# Scale logits and mask invalid
l = l * sm_scale
l = tl.where(valid, l, -float("inf"))
# Streaming softmax update for a single token
m_new = tl.maximum(m, l)
scale_prev = tl.exp(m - m_new)
p = tl.exp(l - m_new)
# Update denominator
S = S * scale_prev + p
# Update numerators
O0 = O0 * scale_prev + K0 * p
O1 = O1 * scale_prev + K1 * p
O2 = O2 * scale_prev + K2 * p
O3 = O3 * scale_prev + K3 * p
m = m_new
start += BLOCK_TOK
inv_S = 1.0 / S
O0 = O0 * inv_S
O1 = O1 * inv_S
O2 = O2 * inv_S
O3 = O3 * inv_S
# Store output
tl.store(out_base + offs_dckv * stride_out_d, O0.to(tl.bfloat16), mask=offs_dckv < DCKV)
tl.store(out_base + d1 * stride_out_d, O1.to(tl.bfloat16), mask=d1 < DCKV)
tl.store(out_base + d2 * stride_out_d, O2.to(tl.bfloat16), mask=d2 < DCKV)
tl.store(out_base + d3 * stride_out_d, O3.to(tl.bfloat16), mask=d3 < DCKV)
# Base-2 LSE: logsumexp(logits_scaled) / log(2)
ln2 = tl.log(2.0)
lse_val = (m + tl.log(S)) / ln2
tl.store(lse_off, lse_val)
def run(q_nope, q_pe, ckv_cache, kpe_cache, kv_indptr, kv_indices, sm_scale):
# Validate dtypes
assert q_nope.dtype == torch.bfloat16, "q_nope must be bfloat16"
assert q_pe.dtype == torch.bfloat16, "q_pe must be bfloat16"
assert ckv_cache.dtype == torch.bfloat16, "ckv_cache must be bfloat16"
assert kpe_cache.dtype == torch.bfloat16, "kpe_cache must be bfloat16"
assert kv_indptr.dtype == torch.int32, "kv_indptr must be int32"
assert kv_indices.dtype == torch.int32, "kv_indices must be int32"
# Shapes and constants
batch_size, num_qo_heads, head_dim_ckv = q_nope.shape
head_dim_kpe = q_pe.shape[-1]
page_size = ckv_cache.shape[1]
len_indptr = kv_indptr.shape[0]
num_kv_indices = kv_indices.shape[0]
assert num_qo_heads == 16, "num_qo_heads must be 16"
assert head_dim_ckv == 512, "head_dim_ckv must be 512"
assert head_dim_kpe == 64, "head_dim_kpe must be 64"
assert page_size == 1, "page_size must be 1"
assert len_indptr == batch_size + 1, "len_indptr must equal batch_size + 1"
assert num_kv_indices == int(kv_indptr[-1].item()), "num_kv_indices must equal kv_indptr[-1]"
# Device handling
orig_device = q_nope.device
if q_nope.is_cuda:
device = q_nope.device
else:
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available, but Triton kernel requires a GPU.")
device = torch.device("cuda")
# Move tensors to device
def to_dev(t):
return t.to(device, non_blocking=True)
q_nope_dev = to_dev(q_nope.contiguous())
q_pe_dev = to_dev(q_pe.contiguous())
# Squeeze page dimension (ps=1)
ckv_dev = to_dev(ckv_cache.squeeze(1).contiguous()) # [num_pages, 512]
kpe_dev = to_dev(kpe_cache.squeeze(1).contiguous()) # [num_pages, 64]
kv_indptr_dev = to_dev(kv_indptr.contiguous())
kv_indices_dev = to_dev(kv_indices.contiguous())
# Outputs
output_dev = torch.empty((batch_size, num_qo_heads, head_dim_ckv), dtype=torch.bfloat16, device=device)
lse_dev = torch.empty((batch_size, num_qo_heads), dtype=torch.float32, device=device)
# Strides (elements)
stride_qn_b, stride_qn_h, stride_qn_d = q_nope_dev.stride()
stride_qp_b, stride_qp_h, stride_qp_d = q_pe_dev.stride()
stride_ckv_p, stride_ckv_d = ckv_dev.stride()
stride_kpe_p, stride_kpe_d = kpe_dev.stride()
stride_out_b, stride_out_h, stride_out_d = output_dev.stride()
stride_lse_b, stride_lse_h = lse_dev.stride()
# Launch configuration
B = batch_size
H = 16
DCKV = 512
DKPE = 64
# Token block and vector block sizes
# Smaller BLOCK_TOK for better numerical stability and register pressure
BLOCK_TOK = 32
BLOCK_DCKV = 128
BLOCK_DKPE = 64
grid = (B * H,)
mla_paged_decode_h16_ckv512_kpe64_ps1_kernel[grid](
q_nope_dev,
q_pe_dev,
ckv_dev,
kpe_dev,
kv_indptr_dev,
kv_indices_dev,
output_dev,
lse_dev,
B,
H,
DCKV,
DKPE,
float(sm_scale),
stride_qn_b,
stride_qn_h,
stride_qn_d,
stride_qp_b,
stride_qp_h,
stride_qp_d,
stride_ckv_p,
stride_ckv_d,
stride_kpe_p,
stride_kpe_d,
stride_out_b,
stride_out_h,
stride_out_d,
stride_lse_b,
stride_lse_h,
BLOCK_TOK,
BLOCK_DCKV,
BLOCK_DKPE,
num_warps=4,
num_stages=2,
)
# Move outputs back to original device
output = output_dev.to(orig_device, non_blocking=True)
lse = lse_dev.to(orig_device, non_blocking=True)
return output, lsescrolls · 286 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON