gemini-2.5-pro / tritonpr9imz
gemini-2.5-pro_triton_pr9imz · gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 264 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-pr9imz?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=15 · num_kv_indices=14
NVIDIA B200
25.0µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=82 · num_kv_indices=65
NVIDIA B200
25.3µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=67 · num_kv_indices=50
NVIDIA B200
25.4µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=18 · num_kv_indices=2
NVIDIA B200
25.5µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=71 · num_kv_indices=54
NVIDIA B200
26.2µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1220 · num_kv_indices=1193
NVIDIA B200
26.3µ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
26.6µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=59 · num_kv_indices=42
NVIDIA B200
26.7µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=9316 · num_kv_indices=73
NVIDIA B200
27.2µ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
27.4µ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=63 · num_kv_indices=46
NVIDIA B200
27.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=55 · num_kv_indices=38
NVIDIA B200
27.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=9341 · num_kv_indices=98
NVIDIA B200
27.9µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=74 · num_kv_indices=57
NVIDIA B200
28.1µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=78 · num_kv_indices=61
NVIDIA B200
29.0µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=406 · num_kv_indices=356
NVIDIA B200
29.5µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=597 · num_kv_indices=547
NVIDIA B200
29.5µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1556 · num_kv_indices=1529
NVIDIA B200
29.8µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2372 · num_kv_indices=2345
NVIDIA B200
31.6µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1069 · num_kv_indices=1034
NVIDIA B200
32.5µ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
32.8µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2052 · num_kv_indices=2025
NVIDIA B200
33.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2548 · num_kv_indices=2521
NVIDIA B200
34.3µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2868 · num_kv_indices=2841
NVIDIA B200
34.8µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=3044 · num_kv_indices=3017
NVIDIA B200
36.9µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1892 · num_kv_indices=1865
NVIDIA B200
42.8µ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
43.2µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=223 · num_kv_indices=173
NVIDIA B200
46.1µ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
66.4µ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
101.4µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=81390 · num_kv_indices=12942
NVIDIA B200
116.2µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=30163 · num_kv_indices=20911
NVIDIA B200
156.4µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=60071 · num_kv_indices=50902
NVIDIA B200
252.7µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=62605 · num_kv_indices=53334
NVIDIA B200
258.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
261.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65421 · num_kv_indices=56150
NVIDIA B200
264.1µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65805 · num_kv_indices=56534
NVIDIA B200
265.1µ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
269.0µ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
270.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67853 · num_kv_indices=58582
NVIDIA B200
274.7µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63053 · num_kv_indices=53782
NVIDIA B200
285.3µ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.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=66253 · num_kv_indices=56982
NVIDIA B200
299.3µ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
328.6µ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
461.9µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=66637 · num_kv_indices=57366
NVIDIA B200
480.9µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63821 · num_kv_indices=54550
NVIDIA B200
893.5µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=64653 · num_kv_indices=55382
NVIDIA B200
903.7µs
#5 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:e5d6283af3f6143bf93e3bc93f551732f18d50a9d53538cf7fead3802f18b330
license declaredApache-2.0
license concludedApache-2.0
authorsgemini-2.5-pro
imported2026-08-20
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
@triton.autotune(num-warps = 4
triton.Config({'BLOCK_KV_LEN': 16}, num_warps=4),Kernel source
main.py264 lines
import torch
import triton
import triton.language as tl
import math
@triton.autotune(
configs=[
triton.Config({'BLOCK_KV_LEN': 16}, num_warps=4),
triton.Config({'BLOCK_KV_LEN': 32}, num_warps=4),
triton.Config({'BLOCK_KV_LEN': 64}, num_warps=4),
triton.Config({'BLOCK_KV_LEN': 128}, num_warps=4),
triton.Config({'BLOCK_KV_LEN': 256}, num_warps=4),
triton.Config({'BLOCK_KV_LEN': 16}, num_warps=8),
triton.Config({'BLOCK_KV_LEN': 32}, num_warps=8),
triton.Config({'BLOCK_KV_LEN': 64}, num_warps=8),
triton.Config({'BLOCK_KV_LEN': 128}, num_warps=8),
triton.Config({'BLOCK_KV_LEN': 256}, num_warps=8),
],
key=['HEAD_DIM'],
)
@triton.jit
def gqa_paged_decode_h32_kv8_d128_ps1_kernel(
# Pointers to Tensors
Q_ptr, K_cache_ptr, V_cache_ptr,
kv_indptr_ptr, kv_indices_ptr,
sm_scale,
Output_ptr, LSE_ptr,
# Stride Info
stride_q_bs, stride_q_h,
stride_k_num_pages, stride_k_ps, stride_k_h,
stride_v_num_pages, stride_v_ps, stride_v_h,
stride_o_bs, stride_o_h,
stride_lse_bs,
# Compile-time Constants
NUM_QO_HEADS: tl.constexpr,
NUM_KV_HEADS: tl.constexpr,
HEAD_DIM: tl.constexpr,
GQA_RATIO: tl.constexpr,
PAGE_SIZE: tl.constexpr,
BLOCK_D: tl.constexpr,
BLOCK_KV_LEN: tl.constexpr,
):
"""
Triton kernel for GQA paged decode attention.
This kernel computes attention for a batch of query vectors against their
corresponding key/value history, which is stored in a paged cache. Each
program instance handles one query head for one sequence in the batch.
Grid: (batch_size, num_qo_heads)
- pid_b (program_id 0): batch index
- pid_h (program_id 1): query head index
Key Optimizations for B200:
- Processes the variable-length KV sequence in fixed-size blocks (BLOCK_KV_LEN)
to increase arithmetic intensity and hide memory latency from gathers.
- Online softmax algorithm is used to compute attention scores and output
in a single pass over the KV sequence, avoiding materialization of the
full attention matrix.
- For this decode-style kernel (one query vector per program), the Q@K.T and P@V
operations are GEMV-like. Since tl.dot has minimum shape requirements (e.g., M>=16)
for Tensor Core usage that are not met by a single vector, these operations
are implemented using efficient element-wise operations and reductions.
- Autotuning is enabled for BLOCK_KV_LEN and num_warps to find the optimal
configuration for the target hardware.
"""
# 1. Get program IDs for batch and query head
pid_b = tl.program_id(0)
pid_h = tl.program_id(1)
# 2. Determine KV sequence length and handle empty sequences
page_start = tl.load(kv_indptr_ptr + pid_b)
page_end = tl.load(kv_indptr_ptr + pid_b + 1)
seq_len = page_end - page_start
# Early exit for sequences with no KV history
if seq_len == 0:
offs_d = tl.arange(0, BLOCK_D)
out_ptr = Output_ptr + pid_b * stride_o_bs + pid_h * stride_o_h + offs_d
lse_ptr = LSE_ptr + pid_b * stride_lse_bs + pid_h
tl.store(out_ptr, tl.zeros([BLOCK_D], dtype=tl.bfloat16), mask=offs_d < BLOCK_D)
tl.store(lse_ptr, -float('inf'))
return
# 3. Load query vector
offs_d = tl.arange(0, BLOCK_D)
q_ptr = Q_ptr + pid_b * stride_q_bs + pid_h * stride_q_h + offs_d
q = tl.load(q_ptr, mask=offs_d < BLOCK_D).to(tl.float32)[None, :] # Shape: [1, BLOCK_D]
# 4. Initialize accumulators for online softmax
# FIX: Use scalar accumulators for m_i and l_i to avoid shape errors
# during the final tl.store operation for the LSE scalar.
acc_o = tl.zeros([BLOCK_D], dtype=tl.float32)
m_i = -float('inf')
l_i = 0.0
# 5. Determine the corresponding KV head for GQA
kv_head_idx = pid_h // GQA_RATIO
# 6. Loop over the KV sequence in blocks
for offset in range(0, seq_len, BLOCK_KV_LEN):
# a. Create masks and pointers for the current block
offs_kv_block = offset + tl.arange(0, BLOCK_KV_LEN)
mask_kv_block = offs_kv_block < seq_len
indices_ptr = kv_indices_ptr + page_start + offs_kv_block
# b. Gather page indices for K and V caches
page_indices = tl.load(indices_ptr, mask=mask_kv_block, other=0)
# c. Gather K vectors for the block
offs_k_h = kv_head_idx * stride_k_h
k_ptrs = K_cache_ptr + page_indices[:, None] * stride_k_num_pages + offs_k_h + offs_d[None, :]
k = tl.load(k_ptrs, mask=mask_kv_block[:, None] & (offs_d[None, :] < BLOCK_D), other=0.0)
# d. Compute scores S = Q @ K.T
s_block = tl.sum(q * k.to(tl.float32), axis=1) * sm_scale
# FIX: Keep scores as a 1D tensor [BLOCK_KV_LEN] for scalar reduction.
s = tl.where(mask_kv_block, s_block, -float('inf'))
# e. Online softmax update (with scalar state)
m_block_max = tl.max(s, axis=0)
m_curr = tl.maximum(m_i, m_block_max)
p = tl.exp(s - m_curr)
l_i_exp = tl.exp(m_i - m_curr)
l_curr = l_i_exp * l_i + tl.sum(p, axis=0)
# f. Gather V vectors for the block
offs_v_h = kv_head_idx * stride_v_h
v_ptrs = V_cache_ptr + page_indices[:, None] * stride_v_num_pages + offs_v_h + offs_d[None, :]
v = tl.load(v_ptrs, mask=mask_kv_block[:, None] & (offs_d[None, :] < BLOCK_D), other=0.0)
# g. Update output accumulator
acc_o = acc_o * l_i_exp
# FIX: Reshape 1D p to [BLOCK_KV_LEN, 1] for broadcasted matmul-like update.
update_o = tl.sum(p[:, None] * v.to(tl.float32), axis=0)
acc_o += update_o
# h. Update state for the next iteration
m_i = m_curr
l_i = l_curr
# 7. Finalize and store results
l_i_safe = tl.where(l_i == 0, 1.0, l_i)
acc_o = acc_o / l_i_safe
# Calculate 2-based log-sum-exp
lse = m_i + tl.log(l_i_safe)
lse = lse / 0.6931471805599453
# Store output vector and LSE value
out_ptr = Output_ptr + pid_b * stride_o_bs + pid_h * stride_o_h + offs_d
tl.store(out_ptr, acc_o.to(tl.bfloat16), mask=offs_d < BLOCK_D)
# FIX: Storing the scalar `lse` value to a scalar pointer is now valid.
tl.store(LSE_ptr + pid_b * stride_lse_bs + pid_h, lse)
def _gqa_paged_decode_h32_kv8_d128_ps1_launcher(q, k_cache, v_cache, kv_indptr, kv_indices, sm_scale):
"""
Host-side wrapper for the Triton kernel.
This function handles device management, output tensor allocation, grid
computation, and kernel invocation. It ensures all tensors are on the
same CUDA device before launching the kernel and moves the results back
to the original device of the `q` tensor.
"""
# 1. Validate inputs and extract dimensions
assert q.dtype == torch.bfloat16
assert k_cache.dtype == torch.bfloat16
assert v_cache.dtype == torch.bfloat16
assert kv_indptr.dtype == torch.int32
assert kv_indices.dtype == torch.int32
batch_size, num_qo_heads, head_dim = q.shape
num_pages, page_size, num_kv_heads, _ = k_cache.shape
# Check problem-specific constants
assert num_qo_heads == 32, f"Expected num_qo_heads=32, got {num_qo_heads}"
assert num_kv_heads == 8, f"Expected num_kv_heads=8, got {num_kv_heads}"
assert head_dim == 128, f"Expected head_dim=128, got {head_dim}"
assert page_size == 1, f"Expected page_size=1, got {page_size}"
# 2. Set default sm_scale if not provided
if sm_scale is None:
sm_scale = 1.0 / math.sqrt(head_dim)
# 3. Complete device management
original_q_device = q.device
# Find a common CUDA device for computation
target_device = None
for t in [q, k_cache, v_cache, kv_indptr, kv_indices]:
if t.is_cuda:
target_device = t.device
break
if target_device is None:
if torch.cuda.is_available():
target_device = torch.device("cuda")
else:
raise RuntimeError("This Triton kernel requires a CUDA-enabled GPU, but none was found.")
# Move all tensors to the target CUDA device
q = q.to(target_device)
k_cache = k_cache.to(target_device)
v_cache = v_cache.to(target_device)
kv_indptr = kv_indptr.to(target_device)
kv_indices = kv_indices.to(target_device)
# 4. Allocate output tensors on the target device
output = torch.empty_like(q)
lse = torch.empty((batch_size, num_qo_heads), dtype=torch.float32, device=target_device)
# 5. Set up grid and launch kernel
grid = (batch_size, num_qo_heads)
gqa_paged_decode_h32_kv8_d128_ps1_kernel[grid](
q, k_cache, v_cache,
kv_indptr, kv_indices,
float(sm_scale),
output, lse,
# Strides
q.stride(0), q.stride(1),
k_cache.stride(0), k_cache.stride(1), k_cache.stride(2),
v_cache.stride(0), v_cache.stride(1), v_cache.stride(2),
output.stride(0), output.stride(1),
lse.stride(0),
# Constants
NUM_QO_HEADS=num_qo_heads,
NUM_KV_HEADS=num_kv_heads,
HEAD_DIM=head_dim,
GQA_RATIO=num_qo_heads // num_kv_heads,
PAGE_SIZE=page_size,
BLOCK_D=head_dim,
)
# 6. Move results back to the original device of `q`
output = output.to(original_q_device)
lse = lse.to(original_q_device)
return output, lse
def run(*args, **kwargs):
"""
Public entry point for the GQA paged decode attention kernel.
This function acts as a flexible interface, accepting both positional and
keyword arguments and forwarding them to the core launcher function.
Args:
q (torch.Tensor): Query tensor of shape [batch_size, 32, 128] and dtype bfloat16.
k_cache (torch.Tensor): Key cache of shape [num_pages, 1, 8, 128] and dtype bfloat16.
v_cache (torch.Tensor): Value cache of shape [num_pages, 1, 8, 128] and dtype bfloat16.
kv_indptr (torch.Tensor): KV page offsets of shape [batch_size + 1] and dtype int32.
kv_indices (torch.Tensor): Page IDs of shape [num_kv_indices] and dtype int32.
sm_scale (float, optional): Softmax scale. Defaults to 1/sqrt(head_dim).
Returns:
Tuple[torch.Tensor, torch.Tensor]:
- output: The attention output tensor of shape [batch_size, 32, 128] and dtype bfloat16.
- lse: The log-sum-exp of attention logits (base 2) of shape [batch_size, 32] and dtype float32.
"""
return _gqa_paged_decode_h32_kv8_d128_ps1_launcher(*args, **kwargs)
scrolls · 264 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON