gpt-o3 / tritonc3c0cc
gpt-o3_triton_c3c0cc · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 164 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-c3c0cc?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32
Benchmark evidence
48 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=18 · num_kv_indices=2
NVIDIA B200
10.3µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=15 · num_kv_indices=14
NVIDIA B200
10.4µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=11 · num_kv_indices=10
NVIDIA B200
10.4µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=55 · num_kv_indices=38
NVIDIA B200
11.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=63 · num_kv_indices=46
NVIDIA B200
12.3µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=67 · num_kv_indices=50
NVIDIA B200
12.3µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=59 · num_kv_indices=42
NVIDIA B200
12.3µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=71 · num_kv_indices=54
NVIDIA B200
12.3µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=74 · num_kv_indices=57
NVIDIA B200
12.4µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=78 · num_kv_indices=61
NVIDIA B200
12.5µs
#2 of 7
2025-10-16
Show all 48 measurements ›Showing all 48 measurements ⌄
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=82 · num_kv_indices=65
NVIDIA B200
12.5µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=9316 · num_kv_indices=73
NVIDIA B200
12.7µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=9341 · num_kv_indices=98
NVIDIA B200
14.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=223 · num_kv_indices=173
NVIDIA B200
19.9µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=406 · num_kv_indices=356
NVIDIA B200
28.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1069 · num_kv_indices=1034
NVIDIA B200
30.6µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1220 · num_kv_indices=1193
NVIDIA B200
30.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1556 · num_kv_indices=1529
NVIDIA B200
30.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1396 · num_kv_indices=1369
NVIDIA B200
31.0µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1732 · num_kv_indices=1705
NVIDIA B200
32.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1892 · num_kv_indices=1865
NVIDIA B200
35.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2052 · num_kv_indices=2025
NVIDIA B200
35.9µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2212 · num_kv_indices=2185
NVIDIA B200
36.7µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2372 · num_kv_indices=2345
NVIDIA B200
37.0µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2548 · num_kv_indices=2521
NVIDIA B200
38.4µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2708 · num_kv_indices=2681
NVIDIA B200
38.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2868 · num_kv_indices=2841
NVIDIA B200
39.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=3044 · num_kv_indices=3017
NVIDIA B200
39.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=597 · num_kv_indices=547
NVIDIA B200
41.1µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=4333 · num_kv_indices=4298
NVIDIA B200
46.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=81390 · num_kv_indices=12942
NVIDIA B200
169.7µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=30163 · num_kv_indices=20911
NVIDIA B200
204.4µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=60071 · num_kv_indices=50902
NVIDIA B200
277.9µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63053 · num_kv_indices=53782
NVIDIA B200
288.2µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=62605 · num_kv_indices=53334
NVIDIA B200
289.0µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63437 · num_kv_indices=54166
NVIDIA B200
291.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=64205 · num_kv_indices=54934
NVIDIA B200
292.6µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63821 · num_kv_indices=54550
NVIDIA B200
292.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65037 · num_kv_indices=55766
NVIDIA B200
293.3µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65805 · num_kv_indices=56534
NVIDIA B200
295.0µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65421 · num_kv_indices=56150
NVIDIA B200
295.4µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=64653 · num_kv_indices=55382
NVIDIA B200
295.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67021 · num_kv_indices=57750
NVIDIA B200
297.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=66253 · num_kv_indices=56982
NVIDIA B200
298.8µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=66637 · num_kv_indices=57366
NVIDIA B200
299.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67405 · num_kv_indices=58134
NVIDIA B200
300.1µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67853 · num_kv_indices=58582
NVIDIA B200
301.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=68237 · num_kv_indices=58966
NVIDIA B200
309.0µs
#3 of 7
2025-10-16
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:55702029eab3dc4e109a991973a4c7a8b689d1790d073c27b8e9a8ad1c7a22d9
license declaredApache-2.0
license concludedApache-2.0
authorsgpt-o3
imported2026-08-20
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
num_warps=4,stages = 4
num_stages=4,Kernel source
main.py164 lines
import math
import torch
import triton
import triton.language as tl
@triton.jit
def gqa_paged_decode_kernel(
q_ptr, # *bf16 [B, 32, 128]
k_ptr, # *bf16 [N_pages, 8, 128] (page_size squeezed)
v_ptr, # *bf16 [N_pages, 8, 128] (page_size squeezed)
kv_indptr_ptr, # *int32 [B + 1]
kv_indices_ptr, # *int32 [num_kv_indices]
sm_scale, # fp32 scalar
out_ptr, # *bf16 [B, 32, 128]
lse_ptr, # *fp32 [B, 32]
BLOCK_T: tl.constexpr,
HEAD_DIM: tl.constexpr,
NUM_QO_HEADS: tl.constexpr,
NUM_KV_HEADS: tl.constexpr,
):
pid = tl.program_id(0)
batch_idx = pid // NUM_QO_HEADS
qo_head = pid % NUM_QO_HEADS
gqa_ratio = NUM_QO_HEADS // NUM_KV_HEADS
kv_head = qo_head // gqa_ratio
# ---- strides (in elements, not bytes) ----
stride_q_batch = NUM_QO_HEADS * HEAD_DIM
stride_q_head = HEAD_DIM
stride_k_page = NUM_KV_HEADS * HEAD_DIM # page_size = 1
stride_k_kv_head = HEAD_DIM
stride_v_page = stride_k_page
stride_v_kv_head = HEAD_DIM
# ---- load query vector ----
d_offs = tl.arange(0, HEAD_DIM)
q_ptr_head = q_ptr + batch_idx * stride_q_batch + qo_head * stride_q_head + d_offs
q_vec = tl.cast(tl.load(q_ptr_head), tl.float32)
# ---- sequence token range ----
start = tl.load(kv_indptr_ptr + batch_idx)
end = tl.load(kv_indptr_ptr + batch_idx + 1)
num_tokens = end - start
# ---- streaming softmax vars ----
m_val = tl.full([], -1e30, tl.float32) # running max
d_val = tl.zeros([], tl.float32) # running sum exp
o_vec = tl.zeros([HEAD_DIM], tl.float32) # running output vector
offset = tl.zeros([], tl.int32)
while offset < num_tokens:
t_offs = tl.arange(0, BLOCK_T)
remain = num_tokens - offset
tok_mask = t_offs < remain
# ---- load page indices ----
pages = tl.load(kv_indices_ptr + start + offset + t_offs,
mask=tok_mask, other=0)
# ---- gather K / V ----
k_ptrs = k_ptr + pages[:, None] * stride_k_page + kv_head * stride_k_kv_head + d_offs[None, :]
v_ptrs = v_ptr + pages[:, None] * stride_v_page + kv_head * stride_v_kv_head + d_offs[None, :]
k_block = tl.cast(tl.load(k_ptrs, mask=tok_mask[:, None], other=0), tl.float32)
v_block = tl.cast(tl.load(v_ptrs, mask=tok_mask[:, None], other=0), tl.float32)
# ---- logits ----
logits = tl.sum(k_block * q_vec[None, :], axis=1) * sm_scale
logits = tl.where(tok_mask, logits, -1e30)
# ---- block softmax ----
m_block = tl.max(logits, axis=0)
exp_logits = tl.exp(logits - m_block)
sum_exp_block = tl.sum(exp_logits, axis=0)
weighted_v = tl.sum(exp_logits[:, None] * v_block, axis=0)
# ---- merge with running values ----
new_m = tl.maximum(m_val, m_block)
alpha_prev = tl.exp(m_val - new_m)
alpha_blk = tl.exp(m_block - new_m)
o_vec = o_vec * alpha_prev + weighted_v * alpha_blk
d_val = d_val * alpha_prev + sum_exp_block * alpha_blk
m_val = new_m
offset += BLOCK_T
inv_d = tl.where(d_val == 0, 0.0, 1.0 / d_val)
out_vec = o_vec * inv_d
log2e = 1.4426950408889634
lse_val = tl.where(d_val == 0,
-1e30,
(tl.log(d_val) + m_val) * log2e)
# ---- store ----
out_ptr_head = out_ptr + batch_idx * stride_q_batch + qo_head * stride_q_head + d_offs
tl.store(out_ptr_head, tl.cast(out_vec, tl.bfloat16))
lse_ptr_head = lse_ptr + batch_idx * NUM_QO_HEADS + qo_head
tl.store(lse_ptr_head, lse_val)
def run(q,
k_cache,
v_cache,
kv_indptr,
kv_indices,
sm_scale: float | None = None):
"""
Entry point for gqa_paged_decode_h32_kv8_d128_ps1.
Returns (output, lse).
"""
if sm_scale is None:
sm_scale = 1.0 / math.sqrt(128.0)
if not torch.cuda.is_available():
raise RuntimeError("CUDA device is required to run Triton kernels.")
# move tensors to GPU if necessary
tensors = [q, k_cache, v_cache, kv_indptr, kv_indices]
device_tensors = [t.cuda() if not t.is_cuda else t for t in tensors]
q_dev, k_dev, v_dev, iptr_dev, idx_dev = [t.contiguous() for t in device_tensors]
batch_size = q_dev.shape[0]
num_qo_heads = 32
head_dim = 128
# squeeze page dimension (=1)
k_dev_flat = k_dev.squeeze(1).contiguous()
v_dev_flat = v_dev.squeeze(1).contiguous()
out_dev = torch.empty((batch_size, num_qo_heads, head_dim),
dtype=torch.bfloat16,
device=q_dev.device)
lse_dev = torch.empty((batch_size, num_qo_heads),
dtype=torch.float32,
device=q_dev.device)
# launch kernel
BLOCK_T = 128
grid = (batch_size * num_qo_heads,)
gqa_paged_decode_kernel[grid](
q_dev, k_dev_flat, v_dev_flat,
iptr_dev, idx_dev,
sm_scale,
out_dev, lse_dev,
BLOCK_T=BLOCK_T,
HEAD_DIM=128,
NUM_QO_HEADS=32,
NUM_KV_HEADS=8,
num_warps=4,
num_stages=4,
)
# move back to original device if needed
if not q.is_cuda:
return out_dev.cpu(), lse_dev.cpu()
return out_dev, lse_devscrolls · 164 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON