Skip to content
KernelIndex
Search⌘K

gpt-5_triton_f88811

gpt-5-2025-08-07 · triton · Apache-2.0

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 ›
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 = 8num_warps=8,
online-softmaxm_new = tl.maximum(m_i, m_curr)
stages = 2num_stages=2,
tile-m = 128BLOCK_M = 128 # tokens per block

Kernel 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, lse
scrolls · 282 lines total

Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0

Best evidence level for this revision: reproducible

JSON