gpt-o3 / tritonb12b97
gpt-o3_triton_b12b97 · gpt-o3 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 253 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gpt-o3-triton-b12b97?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32
Benchmark evidence
92 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=15 · num_kv_indices=14
NVIDIA B200
86.6µs
#6 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=81 · num_kv_indices=64
NVIDIA B200
87.0µs
#4 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9332 · num_kv_indices=87
NVIDIA B200
87.1µs
#2 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=8 · num_kv_indices=7
NVIDIA B200
87.7µs
#6 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=17 · num_kv_indices=2
NVIDIA B200
87.8µs
#5 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=10 · num_kv_indices=9
NVIDIA B200
88.2µs
#5 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=12 · num_kv_indices=11
NVIDIA B200
88.7µs
#5 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=191 · num_kv_indices=141
NVIDIA B200
89.3µs
#3 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9317 · num_kv_indices=72
NVIDIA B200
89.4µs
#3 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=57 · num_kv_indices=40
NVIDIA B200
90.3µs
#6 of 6
2026-03-28
Show all 92 measurements ›Showing all 92 measurements ⌄
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=8 · num_kv_indices=7
NVIDIA B200
92.6µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=81 · num_kv_indices=64
NVIDIA B200
93.1µ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
93.6µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=15 · num_kv_indices=14
NVIDIA B200
93.7µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=302 · num_kv_indices=252
NVIDIA B200
94.7µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=57 · num_kv_indices=40
NVIDIA B200
95.5µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=10 · num_kv_indices=9
NVIDIA B200
96.5µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9317 · num_kv_indices=72
NVIDIA B200
96.6µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=67 · num_kv_indices=50
NVIDIA B200
96.8µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=12 · num_kv_indices=11
NVIDIA B200
97.0µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=191 · num_kv_indices=141
NVIDIA B200
98.3µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9347 · num_kv_indices=102
NVIDIA B200
99.4µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=412 · num_kv_indices=362
NVIDIA B200
99.7µs
#2 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9332 · num_kv_indices=87
NVIDIA B200
104.3µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=412 · num_kv_indices=362
NVIDIA B200
106.3µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=596 · num_kv_indices=546
NVIDIA B200
115.0µs
#3 of 6
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=596 · num_kv_indices=546
NVIDIA B200
116.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=486 · num_kv_indices=436
NVIDIA B200
119.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=1070 · num_kv_indices=1020
NVIDIA B200
779.6µs
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=2158 · num_kv_indices=2108
NVIDIA B200
836.2µs
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=1070 · num_kv_indices=1020
NVIDIA B200
861.8µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=3246 · num_kv_indices=3196
NVIDIA B200
964.8µs
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=2158 · num_kv_indices=2108
NVIDIA B200
965.6µs
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=4334 · num_kv_indices=4284
NVIDIA B200
996.1µs
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=4334 · num_kv_indices=4284
NVIDIA B200
1.06ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=3246 · num_kv_indices=3196
NVIDIA B200
1.07ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=5422 · num_kv_indices=5372
NVIDIA B200
1.15ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=5422 · num_kv_indices=5372
NVIDIA B200
1.15ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=6510 · num_kv_indices=6460
NVIDIA B200
1.18ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=7598 · num_kv_indices=7548
NVIDIA B200
1.20ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=8686 · num_kv_indices=8636
NVIDIA B200
1.24ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=7598 · num_kv_indices=7548
NVIDIA B200
1.28ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=8686 · num_kv_indices=8636
NVIDIA B200
1.30ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=6510 · num_kv_indices=6460
NVIDIA B200
1.43ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=24732 · num_kv_indices=15463
NVIDIA B200
1.66ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=25820 · num_kv_indices=16551
NVIDIA B200
1.74ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=26908 · num_kv_indices=17639
NVIDIA B200
1.77ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=24732 · num_kv_indices=15463
NVIDIA B200
1.78ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=25820 · num_kv_indices=16551
NVIDIA B200
1.79ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=27996 · num_kv_indices=18727
NVIDIA B200
1.88ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=29084 · num_kv_indices=19815
NVIDIA B200
1.91ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=26908 · num_kv_indices=17639
NVIDIA B200
1.91ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=29084 · num_kv_indices=19815
NVIDIA B200
1.96ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=27996 · num_kv_indices=18727
NVIDIA B200
1.99ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=31260 · num_kv_indices=21991
NVIDIA B200
2.06ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=30172 · num_kv_indices=20903
NVIDIA B200
2.09ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=30172 · num_kv_indices=20903
NVIDIA B200
2.10ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=31260 · num_kv_indices=21991
NVIDIA B200
2.11ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=32348 · num_kv_indices=23079
NVIDIA B200
2.14ms
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=32348 · num_kv_indices=23079
NVIDIA B200
2.22ms
#2 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=28831 · num_kv_indices=28815
NVIDIA B200
4.45ms
#5 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=31007 · num_kv_indices=30991
NVIDIA B200
4.57ms
#5 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=33183 · num_kv_indices=33167
NVIDIA B200
4.74ms
#5 of 5
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=28831 · num_kv_indices=28815
NVIDIA B200
4.80ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=35359 · num_kv_indices=35343
NVIDIA B200
4.94ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=31007 · num_kv_indices=30991
NVIDIA B200
5.03ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=33183 · num_kv_indices=33167
NVIDIA B200
5.06ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=37535 · num_kv_indices=37519
NVIDIA B200
5.14ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=39711 · num_kv_indices=39695
NVIDIA B200
5.21ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=35359 · num_kv_indices=35343
NVIDIA B200
5.25ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=41887 · num_kv_indices=41871
NVIDIA B200
5.37ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=37535 · num_kv_indices=37519
NVIDIA B200
5.42ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=41887 · num_kv_indices=41871
NVIDIA B200
5.55ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=39711 · num_kv_indices=39695
NVIDIA B200
5.67ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=44063 · num_kv_indices=44047
NVIDIA B200
5.67ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=44063 · num_kv_indices=44047
NVIDIA B200
5.67ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=46303 · num_kv_indices=46287
NVIDIA B200
5.78ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=48479 · num_kv_indices=48463
NVIDIA B200
5.96ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=46303 · num_kv_indices=46287
NVIDIA B200
5.96ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=50655 · num_kv_indices=50639
NVIDIA B200
6.04ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=48479 · num_kv_indices=48463
NVIDIA B200
6.08ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=55007 · num_kv_indices=54991
NVIDIA B200
6.18ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=57183 · num_kv_indices=57167
NVIDIA B200
6.25ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=52831 · num_kv_indices=52815
NVIDIA B200
6.33ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=55007 · num_kv_indices=54991
NVIDIA B200
6.38ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=59359 · num_kv_indices=59343
NVIDIA B200
6.40ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=59359 · num_kv_indices=59343
NVIDIA B200
6.42ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=52831 · num_kv_indices=52815
NVIDIA B200
6.46ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=61535 · num_kv_indices=61519
NVIDIA B200
6.73ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=57183 · num_kv_indices=57167
NVIDIA B200
6.74ms
#7 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=61535 · num_kv_indices=61519
NVIDIA B200
6.84ms
#4 of 4
2026-03-28
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=50655 · num_kv_indices=50639
NVIDIA B200
7.54ms
#7 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:806a80269911f1ea48e10e81d53090bb50b42e8db650c6eebac66fc60534ed90
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 = 8
num_warps = 8,stages = 4
num_stages = 4,Kernel source
main.py253 lines
import math
import torch
import triton
import triton.language as tl
# ----------------------------------------------------------------------------- #
# Triton Kernel #
# ----------------------------------------------------------------------------- #
@triton.jit
def gqa_paged_decode_kernel(
q_ptr, # *bf16 – [num_qo_heads, head_dim]
k_cache_ptr, # *bf16 – [num_pages, num_kv_heads, head_dim]
v_cache_ptr, # *bf16 – [num_pages, num_kv_heads, head_dim]
kv_indices_ptr, # *i32 – [seq_len]
out_ptr, # *bf16 – [num_qo_heads, head_dim]
lse_ptr, # *fp32 – [num_qo_heads]
sm_scale: tl.float32, # soft-max scale
seq_len: tl.constexpr, # number of tokens for this sequence
stride_q_head: tl.constexpr,
stride_k_page: tl.constexpr,
stride_k_head: tl.constexpr,
stride_v_page: tl.constexpr,
stride_v_head: tl.constexpr,
stride_out_head: tl.constexpr,
stride_lse_head: tl.constexpr,
gqa_ratio: tl.constexpr, # 8 for 32/4
BLOCK_TOKENS: tl.constexpr, # 128
BLOCK_D: tl.constexpr, # 128
):
# --------------------------------------------------------------------- #
# Each program instance processes one query/output (qo) head.
# --------------------------------------------------------------------- #
pid = tl.program_id(axis=0)
h_qo = pid # 0 … 31
h_kv = h_qo // gqa_ratio # corresponding KV head (0 … 3)
d_offsets = tl.arange(0, BLOCK_D)
# --------------------------------------------------------------------- #
# Load query vector (bf16 → fp32)
# --------------------------------------------------------------------- #
q_head_ptr = q_ptr + h_qo * stride_q_head
q_vec = tl.load(q_head_ptr + d_offsets).to(tl.float32) # [D]
# Initial values for online softmax reduction
neg_inf = -1.0e30
m_val = tl.full((), neg_inf, dtype=tl.float32) # running max
# --------------------------------------------------------------------- #
# Pass-1 : find maximum logit
# --------------------------------------------------------------------- #
for t0 in range(0, seq_len, BLOCK_TOKENS):
tok_offsets = tl.arange(0, BLOCK_TOKENS) + t0 # [T]
tok_mask = tok_offsets < seq_len
page_ids = tl.load(kv_indices_ptr + tok_offsets, mask=tok_mask, other=0)
k_ptrs = (
k_cache_ptr
+ page_ids[:, None] * stride_k_page
+ h_kv * stride_k_head
+ d_offsets[None, :]
) # [T, D] pointers
k_block = tl.load(k_ptrs, mask=tok_mask[:, None], other=0).to(tl.float32)
logits = tl.sum(k_block * q_vec[None, :], axis=1) * sm_scale
logits = tl.where(tok_mask, logits, neg_inf)
block_max = tl.max(logits, axis=0)
m_val = tl.maximum(m_val, block_max)
# --------------------------------------------------------------------- #
# Pass-2 : compute lse, softmax, output
# --------------------------------------------------------------------- #
lse_denom = tl.zeros((), tl.float32)
out_acc = tl.zeros((BLOCK_D, ), tl.float32)
for t0 in range(0, seq_len, BLOCK_TOKENS):
tok_offsets = tl.arange(0, BLOCK_TOKENS) + t0
tok_mask = tok_offsets < seq_len
page_ids = tl.load(kv_indices_ptr + tok_offsets, mask=tok_mask, other=0)
k_ptrs = (
k_cache_ptr
+ page_ids[:, None] * stride_k_page
+ h_kv * stride_k_head
+ d_offsets[None, :]
)
v_ptrs = (
v_cache_ptr
+ page_ids[:, None] * stride_v_page
+ h_kv * stride_v_head
+ d_offsets[None, :]
)
k_block = tl.load(k_ptrs, mask=tok_mask[:, None], other=0).to(tl.float32)
v_block = tl.load(v_ptrs, mask=tok_mask[:, None], other=0).to(tl.float32)
logits = tl.sum(k_block * q_vec[None, :], axis=1) * sm_scale
logits = logits - m_val
logits = tl.where(tok_mask, logits, neg_inf)
weights = tl.exp(logits) # [T]
weights = tl.where(tok_mask, weights, 0.0)
lse_denom += tl.sum(weights, axis=0)
out_acc += tl.sum(v_block * weights[:, None], axis=0) # [D]
out_final = out_acc / lse_denom # [D]
# --------------------------------------------------------------------- #
# Write back output and LSE (convert back to bf16 where needed)
# --------------------------------------------------------------------- #
out_head_ptr = out_ptr + h_qo * stride_out_head
tl.store(out_head_ptr + d_offsets, out_final.to(tl.bfloat16))
lse_head_ptr = lse_ptr + h_qo * stride_lse_head
INV_LN2 = 1.4426950408889634 # 1 / ln(2)
l_val = (tl.log(lse_denom) + m_val) * INV_LN2
tl.store(lse_head_ptr, l_val)
# ----------------------------------------------------------------------------- #
# Python Wrapper #
# ----------------------------------------------------------------------------- #
@torch.no_grad()
def run(
q: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
kv_indptr: torch.Tensor,
kv_indices: torch.Tensor,
sm_scale: float | None = None,
):
"""
Parameters follow the specification in the problem statement.
Returns
-------
output : bfloat16[batch, 32, 128]
lse : float32[batch, 32]
"""
if sm_scale is None:
sm_scale = 1.0 / math.sqrt(128.0)
# --------------------------------------------------------------------- #
# Device handling
# --------------------------------------------------------------------- #
orig_device = q.device
cuda_available = torch.cuda.is_available()
def to_cuda(t):
if t.device.type == 'cuda':
return t
if not cuda_available:
raise RuntimeError("CUDA is required but not available.")
return t.cuda(non_blocking=True)
q = to_cuda(q).contiguous()
k_cache = to_cuda(k_cache).contiguous()
v_cache = to_cuda(v_cache).contiguous()
kv_indptr = to_cuda(kv_indptr).contiguous()
kv_indices = to_cuda(kv_indices).contiguous()
# --------------------------------------------------------------------- #
# Assertions & shapes
# --------------------------------------------------------------------- #
batch_size, num_qo_heads, head_dim = q.shape
num_pages, page_sz, num_kv_heads, _ = k_cache.shape
assert num_qo_heads == 32, "num_qo_heads must be 32"
assert num_kv_heads == 4, "num_kv_heads must be 4"
assert head_dim == 128
assert page_sz == 1
assert kv_indptr.numel() == batch_size + 1
assert kv_indices.numel() == kv_indptr[-1].item()
# Flatten caches (page_size == 1)
k_cache_flat = k_cache.squeeze(1).contiguous() # [num_pages, 4, 128]
v_cache_flat = v_cache.squeeze(1).contiguous()
# Allocate output tensors
out = torch.zeros_like(q, dtype=torch.bfloat16, device=q.device)
lse = torch.full((batch_size, num_qo_heads),
-float("inf"), dtype=torch.float32, device=q.device)
# Strides (in elements, not bytes)
stride_q_head = head_dim
stride_k_head = head_dim
stride_k_page = num_kv_heads * head_dim # 4 * 128
stride_v_head = head_dim
stride_v_page = num_kv_heads * head_dim
stride_out_head = head_dim
stride_lse_head = 1
gqa_ratio = num_qo_heads // num_kv_heads # 8
BLOCK_TOKENS = 128
BLOCK_D = 128
# --------------------------------------------------------------------- #
# Launch per-sequence kernel
# --------------------------------------------------------------------- #
for b in range(batch_size):
start = kv_indptr[b].item()
end = kv_indptr[b + 1].item()
seq_len = end - start
if seq_len == 0:
continue
seq_indices = kv_indices[start:end].contiguous()
grid = (num_qo_heads,)
gqa_paged_decode_kernel[grid](
q[b],
k_cache_flat,
v_cache_flat,
seq_indices,
out[b],
lse[b],
sm_scale,
seq_len,
stride_q_head,
stride_k_page,
stride_k_head,
stride_v_page,
stride_v_head,
stride_out_head,
stride_lse_head,
gqa_ratio,
BLOCK_TOKENS,
BLOCK_D,
num_warps = 8,
num_stages = 4,
)
# --------------------------------------------------------------------- #
# Move outputs back if the original tensors were on CPU
# --------------------------------------------------------------------- #
if orig_device.type == 'cpu':
out = out.to(orig_device)
lse = lse.to(orig_device)
return out, lsescrolls · 253 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reproducible
JSON