gpt-5_triton_f88811
gpt-5-2025-08-07 · triton · Apache-2.0
Kernel source · 282 lines ↓holds 40 records
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 282 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-5-triton-f88811?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32
Benchmark evidence
96 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=67 · num_kv_indices=50
NVIDIA B200
57.1µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=81 · num_kv_indices=64
NVIDIA B200
57.4µs
#1 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=17 · num_kv_indices=2
NVIDIA B200
57.4µs
#4 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=57 · num_kv_indices=40
NVIDIA B200
57.5µs
#3 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9347 · num_kv_indices=102
NVIDIA B200
57.6µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9332 · num_kv_indices=87
NVIDIA B200
57.6µs
#1 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=15 · num_kv_indices=14
NVIDIA B200
57.8µs
#5 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9317 · num_kv_indices=72
NVIDIA B200
57.9µs
#1 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=12 · num_kv_indices=11
NVIDIA B200
58.6µs
#4 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=10 · num_kv_indices=9
NVIDIA B200
59.0µs
#4 of 5
2026-03-28
Show all 96 measurements ›Showing all 96 measurements ⌄
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=8 · num_kv_indices=7
NVIDIA B200
59.6µs
#5 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=191 · num_kv_indices=141
NVIDIA B200
59.8µs
#2 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=302 · num_kv_indices=252
NVIDIA B200
61.2µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9332 · num_kv_indices=87
NVIDIA B200
65.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=67 · num_kv_indices=50
NVIDIA B200
65.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=12 · num_kv_indices=11
NVIDIA B200
66.1µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=15 · num_kv_indices=14
NVIDIA B200
67.7µs
#6 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=8 · num_kv_indices=7
NVIDIA B200
67.8µs
#6 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9317 · num_kv_indices=72
NVIDIA B200
67.8µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=81 · num_kv_indices=64
NVIDIA B200
67.8µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9347 · num_kv_indices=102
NVIDIA B200
67.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=10 · num_kv_indices=9
NVIDIA B200
68.3µs
#6 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=17 · num_kv_indices=2
NVIDIA B200
68.3µs
#6 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=57 · num_kv_indices=40
NVIDIA B200
69.9µs
#6 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=486 · num_kv_indices=436
NVIDIA B200
71.8µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=412 · num_kv_indices=362
NVIDIA B200
72.2µs
#1 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=191 · num_kv_indices=141
NVIDIA B200
72.4µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=1070 · num_kv_indices=1020
NVIDIA B200
73.8µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=302 · num_kv_indices=252
NVIDIA B200
76.1µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=596 · num_kv_indices=546
NVIDIA B200
80.6µs
#2 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=412 · num_kv_indices=362
NVIDIA B200
81.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=486 · num_kv_indices=436
NVIDIA B200
82.5µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=2158 · num_kv_indices=2108
NVIDIA B200
84.1µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=1070 · num_kv_indices=1020
NVIDIA B200
86.5µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=3246 · num_kv_indices=3196
NVIDIA B200
87.2µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=596 · num_kv_indices=546
NVIDIA B200
88.6µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=2158 · num_kv_indices=2108
NVIDIA B200
93.1µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=4334 · num_kv_indices=4284
NVIDIA B200
97.3µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=5422 · num_kv_indices=5372
NVIDIA B200
101.9µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=3246 · num_kv_indices=3196
NVIDIA B200
102.0µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=4334 · num_kv_indices=4284
NVIDIA B200
105.4µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=6510 · num_kv_indices=6460
NVIDIA B200
109.7µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=5422 · num_kv_indices=5372
NVIDIA B200
116.3µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=7598 · num_kv_indices=7548
NVIDIA B200
117.3µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=6510 · num_kv_indices=6460
NVIDIA B200
119.9µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=8686 · num_kv_indices=8636
NVIDIA B200
122.7µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=7598 · num_kv_indices=7548
NVIDIA B200
128.3µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=8686 · num_kv_indices=8636
NVIDIA B200
136.1µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=24732 · num_kv_indices=15463
NVIDIA B200
199.0µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=25820 · num_kv_indices=16551
NVIDIA B200
207.9µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=26908 · num_kv_indices=17639
NVIDIA B200
207.9µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=27996 · num_kv_indices=18727
NVIDIA B200
212.7µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=29084 · num_kv_indices=19815
NVIDIA B200
216.6µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=24732 · num_kv_indices=15463
NVIDIA B200
219.6µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=30172 · num_kv_indices=20903
NVIDIA B200
221.3µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=28831 · num_kv_indices=28815
NVIDIA B200
226.7µs
#1 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=31260 · num_kv_indices=21991
NVIDIA B200
227.4µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=25820 · num_kv_indices=16551
NVIDIA B200
229.3µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=32348 · num_kv_indices=23079
NVIDIA B200
230.1µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=26908 · num_kv_indices=17639
NVIDIA B200
230.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=27996 · num_kv_indices=18727
NVIDIA B200
232.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=31007 · num_kv_indices=30991
NVIDIA B200
232.5µs
#1 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=29084 · num_kv_indices=19815
NVIDIA B200
235.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=33183 · num_kv_indices=33167
NVIDIA B200
238.6µs
#1 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=28831 · num_kv_indices=28815
NVIDIA B200
240.6µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=30172 · num_kv_indices=20903
NVIDIA B200
245.4µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=35359 · num_kv_indices=35343
NVIDIA B200
248.0µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=31260 · num_kv_indices=21991
NVIDIA B200
250.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=31007 · num_kv_indices=30991
NVIDIA B200
253.8µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=32348 · num_kv_indices=23079
NVIDIA B200
254.9µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=33183 · num_kv_indices=33167
NVIDIA B200
255.0µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=37535 · num_kv_indices=37519
NVIDIA B200
263.7µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=35359 · num_kv_indices=35343
NVIDIA B200
264.9µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=39711 · num_kv_indices=39695
NVIDIA B200
272.1µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=41887 · num_kv_indices=41871
NVIDIA B200
281.9µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=37535 · num_kv_indices=37519
NVIDIA B200
283.5µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=44063 · num_kv_indices=44047
NVIDIA B200
284.3µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=39711 · num_kv_indices=39695
NVIDIA B200
288.4µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=41887 · num_kv_indices=41871
NVIDIA B200
296.8µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=46303 · num_kv_indices=46287
NVIDIA B200
301.1µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=44063 · num_kv_indices=44047
NVIDIA B200
303.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=48479 · num_kv_indices=48463
NVIDIA B200
303.6µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=50655 · num_kv_indices=50639
NVIDIA B200
307.7µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=46303 · num_kv_indices=46287
NVIDIA B200
319.4µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=52831 · num_kv_indices=52815
NVIDIA B200
322.2µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=48479 · num_kv_indices=48463
NVIDIA B200
325.4µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=50655 · num_kv_indices=50639
NVIDIA B200
327.0µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=55007 · num_kv_indices=54991
NVIDIA B200
329.9µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=57183 · num_kv_indices=57167
NVIDIA B200
336.6µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=59359 · num_kv_indices=59343
NVIDIA B200
346.1µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=55007 · num_kv_indices=54991
NVIDIA B200
348.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=57183 · num_kv_indices=57167
NVIDIA B200
355.7µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=61535 · num_kv_indices=61519
NVIDIA B200
356.4µs
#1 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=61535 · num_kv_indices=61519
NVIDIA B200
377.9µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=52831 · num_kv_indices=52815
NVIDIA B200
379.3µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=59359 · num_kv_indices=59343
NVIDIA B200
403.0µs
#2 of 7
2025-10-16
Reproduction-ready · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:09a3a02d28d801076ad976133b410f793223f2428100276618e716be6d2567ac
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.
num-warps = 8
num_warps=8,online-softmax
m_new = tl.maximum(m_i, m_curr)stages = 2
num_stages=2,tile-m = 128
BLOCK_M = 128 # tokens per blockKernel source
main.py282 lines
import math
import torch
import triton
import triton.language as tl
@triton.jit
def gqa_paged_decode_h32_kv4_d128_ps1_kernel(
q_ptr, # *bf16 [B, Hq, D]
k_cache_ptr, # *bf16 [P, S=1, Hk=4, D]
v_cache_ptr, # *bf16 [P, S=1, Hk=4, D]
kv_indptr_ptr, # *i32 [B+1]
kv_indices_ptr, # *i32 [N]
sm_scale, # f32 scalar
output_ptr, # *bf16 [B, Hq, D]
lse_ptr, # *f32 [B, Hq]
batch_size, # i32
num_qo_heads, # i32 (should be 32)
# strides for q
stride_q_b, stride_q_h, stride_q_d,
# strides for k_cache
stride_k_p, stride_k_s, stride_k_h, stride_k_d,
# strides for v_cache
stride_v_p, stride_v_s, stride_v_h, stride_v_d,
# strides for output
stride_o_b, stride_o_h, stride_o_d,
# strides for lse
stride_lse_b, stride_lse_h,
BLOCK_M: tl.constexpr, # token block size
D_HEAD: tl.constexpr, # head dim = 128
GQA_RATIO: tl.constexpr # 8
):
pid = tl.program_id(0)
# Compute (b, h)
h = pid % num_qo_heads
b = pid // num_qo_heads
if b >= batch_size:
return
# Load KV range for this batch
page_start = tl.load(kv_indptr_ptr + b).to(tl.int32)
page_end = tl.load(kv_indptr_ptr + (b + 1)).to(tl.int32)
seq_len = page_end - page_start
# Early exit for empty sequences
if seq_len <= 0:
d_offsets = tl.arange(0, D_HEAD)
o_ptrs = output_ptr + b * stride_o_b + h * stride_o_h + d_offsets * stride_o_d
tl.store(o_ptrs, tl.zeros((D_HEAD,), dtype=tl.bfloat16))
lse_out_ptr = lse_ptr + b * stride_lse_b + h * stride_lse_h
neg_inf = -float("inf")
tl.store(lse_out_ptr, tl.full((), neg_inf, dtype=tl.float32))
return
# Load Q vector for this (b, h)
d_offsets = tl.arange(0, D_HEAD)
q_ptrs = q_ptr + b * stride_q_b + h * stride_q_h + d_offsets * stride_q_d
q_vec = tl.load(q_ptrs).to(tl.float32)
# Initialize streaming softmax state
neg_inf = tl.full((), -float("inf"), dtype=tl.float32)
m_i = neg_inf
l_i = tl.zeros((), dtype=tl.float32)
acc = tl.zeros((D_HEAD,), dtype=tl.float32)
# Determine KV head from GQA mapping
kv_head = (h // GQA_RATIO).to(tl.int32)
# Iterate over tokens in blocks
pos = tl.zeros((), dtype=tl.int32)
while pos < seq_len:
t_offsets = tl.arange(0, BLOCK_M)
curr = pos + t_offsets
mask_t = curr < seq_len
# Gather page ids
page_ids = tl.load(kv_indices_ptr + page_start + curr, mask=mask_t, other=0).to(tl.int32)
# K pointers [BLOCK_M, D_HEAD]
k_ptrs = (
k_cache_ptr
+ page_ids[:, None] * stride_k_p
+ 0 * stride_k_s
+ kv_head * stride_k_h
+ d_offsets[None, :] * stride_k_d
)
k_block = tl.load(k_ptrs, mask=mask_t[:, None], other=0).to(tl.float32)
# Compute logits for this block: [BLOCK_M]
logits = tl.sum(k_block * q_vec[None, :], axis=1)
logits_scaled = logits * sm_scale
logits_scaled = tl.where(mask_t, logits_scaled, neg_inf)
# Block-level max
m_curr = tl.max(logits_scaled, axis=0)
m_new = tl.maximum(m_i, m_curr)
# Compute p = exp(logits - m_new)
p = tl.exp(logits_scaled - m_new)
# sum of p
l_part = tl.sum(p, axis=0)
# Update l_i
l_i = l_i * tl.exp(m_i - m_new) + l_part
# V pointers and weighted accumulation
v_ptrs = (
v_cache_ptr
+ page_ids[:, None] * stride_v_p
+ 0 * stride_v_s
+ kv_head * stride_v_h
+ d_offsets[None, :] * stride_v_d
)
v_block = tl.load(v_ptrs, mask=mask_t[:, None], other=0).to(tl.float32)
weighted = tl.sum(v_block * p[:, None], axis=0)
# Update accumulator and max
acc = acc * tl.exp(m_i - m_new) + weighted
m_i = m_new
pos += BLOCK_M
# Finalize output
nonempty = l_i > 0.0
out_vec = tl.where(nonempty, acc / l_i, tl.zeros((D_HEAD,), dtype=tl.float32))
inv_ln2 = 1.4426950408889634 # 1 / ln(2)
lse_val = tl.where(nonempty, (tl.log(l_i) + m_i) * inv_ln2, neg_inf)
# Store output and lse
o_ptrs = output_ptr + b * stride_o_b + h * stride_o_h + d_offsets * stride_o_d
tl.store(o_ptrs, out_vec.to(tl.bfloat16))
lse_out_ptr = lse_ptr + b * stride_lse_b + h * stride_lse_h
tl.store(lse_out_ptr, lse_val)
def _ensure_cuda(t: torch.Tensor, device: torch.device):
if t.device.type == "cuda":
if t.device != device:
return t.to(device)
return t
else:
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available, but GPU execution is required.")
return t.to(device, non_blocking=True)
def run(q, k_cache, v_cache, kv_indptr, kv_indices, sm_scale=None):
# Validate and default sm_scale
HEAD_DIM = 128
NUM_QO_HEADS = 32
NUM_KV_HEADS = 4
PAGE_SIZE = 1
GQA_RATIO = NUM_QO_HEADS // NUM_KV_HEADS
if sm_scale is None:
sm_scale = 1.0 / math.sqrt(HEAD_DIM)
# Convert sm_scale to Python float
if isinstance(sm_scale, (float, int)):
sm_scale_val = float(sm_scale)
elif isinstance(sm_scale, torch.Tensor):
if sm_scale.numel() != 1:
raise ValueError("sm_scale must be a scalar.")
sm_scale_val = float(sm_scale.detach().cpu().item())
else:
raise TypeError("sm_scale must be a float, int, or 0-dim torch.Tensor")
# Extract shapes and validate
if q.ndim != 3:
raise ValueError("q must have shape [batch_size, num_qo_heads, head_dim]")
batch_size, num_qo_heads, head_dim = q.shape
if k_cache.ndim != 4:
raise ValueError("k_cache must have shape [num_pages, page_size, num_kv_heads, head_dim]")
if v_cache.ndim != 4:
raise ValueError("v_cache must have shape [num_pages, page_size, num_kv_heads, head_dim]")
num_pages, page_size, num_kv_heads, head_dim_k = k_cache.shape
num_pages_v, page_size_v, num_kv_heads_v, head_dim_v = v_cache.shape
if num_pages != num_pages_v or page_size != page_size_v or num_kv_heads != num_kv_heads_v or head_dim_k != head_dim_v:
raise ValueError("k_cache and v_cache shapes must match")
if num_qo_heads != NUM_QO_HEADS:
raise AssertionError("num_qo_heads must be 32")
if num_kv_heads != NUM_KV_HEADS:
raise AssertionError("num_kv_heads must be 4")
if head_dim != HEAD_DIM or head_dim_k != HEAD_DIM:
raise AssertionError("head_dim must be 128")
if page_size != PAGE_SIZE:
raise AssertionError("page_size must be 1")
if kv_indptr.ndim != 1:
raise ValueError("kv_indptr must be 1-D")
if kv_indices.ndim != 1:
raise ValueError("kv_indices must be 1-D")
len_indptr = kv_indptr.shape[0]
num_kv_indices = kv_indices.shape[0]
if len_indptr != batch_size + 1:
raise AssertionError("len_indptr must be batch_size + 1")
# kv_indptr[-1] value (synchronize to host once)
last_ind = int(kv_indptr[-1].detach().cpu().item())
if num_kv_indices != last_ind:
raise AssertionError("num_kv_indices must equal kv_indptr[-1].item()")
# Dtypes
if q.dtype != torch.bfloat16:
raise TypeError("q must be bfloat16")
if k_cache.dtype != torch.bfloat16 or v_cache.dtype != torch.bfloat16:
raise TypeError("k_cache and v_cache must be bfloat16")
if kv_indptr.dtype != torch.int32 or kv_indices.dtype != torch.int32:
raise TypeError("kv_indptr and kv_indices must be int32")
# Determine working device
if not torch.cuda.is_available():
for t in (q, k_cache, v_cache, kv_indptr, kv_indices):
if t.is_cuda:
raise RuntimeError("Input tensor is on CUDA device, but CUDA is not available.")
raise RuntimeError("CUDA is not available. Triton kernel cannot run on CPU.")
if q.device.type == "cuda":
work_device = q.device
elif k_cache.device.type == "cuda":
work_device = k_cache.device
elif v_cache.device.type == "cuda":
work_device = v_cache.device
elif kv_indptr.device.type == "cuda":
work_device = kv_indptr.device
elif kv_indices.device.type == "cuda":
work_device = kv_indices.device
else:
work_device = torch.device("cuda")
orig_q_device = q.device
# Move tensors to working device
q_dev = _ensure_cuda(q.contiguous(), work_device)
k_cache_dev = _ensure_cuda(k_cache.contiguous(), work_device)
v_cache_dev = _ensure_cuda(v_cache.contiguous(), work_device)
kv_indptr_dev = _ensure_cuda(kv_indptr.contiguous(), work_device)
kv_indices_dev = _ensure_cuda(kv_indices.contiguous(), work_device)
# Allocate outputs on working device
output_dev = torch.empty((batch_size, NUM_QO_HEADS, HEAD_DIM), dtype=torch.bfloat16, device=work_device)
lse_dev = torch.empty((batch_size, NUM_QO_HEADS), dtype=torch.float32, device=work_device)
# Launch kernel
BLOCK_M = 128 # tokens per block
grid = (batch_size * NUM_QO_HEADS,)
gqa_paged_decode_h32_kv4_d128_ps1_kernel[grid](
q_dev,
k_cache_dev,
v_cache_dev,
kv_indptr_dev,
kv_indices_dev,
sm_scale_val,
output_dev,
lse_dev,
batch_size,
NUM_QO_HEADS,
# q strides
q_dev.stride(0), q_dev.stride(1), q_dev.stride(2),
# k strides
k_cache_dev.stride(0), k_cache_dev.stride(1), k_cache_dev.stride(2), k_cache_dev.stride(3),
# v strides
v_cache_dev.stride(0), v_cache_dev.stride(1), v_cache_dev.stride(2), v_cache_dev.stride(3),
# output strides
output_dev.stride(0), output_dev.stride(1), output_dev.stride(2),
# lse strides
lse_dev.stride(0), lse_dev.stride(1),
BLOCK_M=BLOCK_M,
D_HEAD=HEAD_DIM,
GQA_RATIO=GQA_RATIO,
num_warps=8,
num_stages=2,
)
# Move outputs back to original device of q
if orig_q_device != work_device:
output = output_dev.to(orig_q_device)
lse = lse_dev.to(orig_q_device)
else:
output = output_dev
lse = lse_dev
return output, lsescrolls · 282 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reproducible
JSON