gemini-2.5-pro / tritonh7ykt0
gemini-2.5-pro_triton_h7ykt0 · gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 268 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-h7ykt0?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 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9332 · num_kv_indices=87
NVIDIA B200
64.9µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=10 · num_kv_indices=9
NVIDIA B200
66.2µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9317 · num_kv_indices=72
NVIDIA B200
66.6µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=15 · num_kv_indices=14
NVIDIA B200
66.8µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=57 · num_kv_indices=40
NVIDIA B200
66.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=67 · num_kv_indices=50
NVIDIA B200
66.9µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=12 · num_kv_indices=11
NVIDIA B200
67.3µs
#6 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.4µs
#2 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.5µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=17 · num_kv_indices=2
NVIDIA B200
67.5µs
#5 of 7
2025-10-16
Show all 48 measurements ›Showing all 48 measurements ⌄
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9347 · num_kv_indices=102
NVIDIA B200
67.5µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=191 · num_kv_indices=141
NVIDIA B200
78.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=302 · num_kv_indices=252
NVIDIA B200
79.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=1070 · num_kv_indices=1020
NVIDIA B200
93.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=412 · num_kv_indices=362
NVIDIA B200
94.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=2158 · num_kv_indices=2108
NVIDIA B200
103.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=3246 · num_kv_indices=3196
NVIDIA B200
104.3µ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
106.7µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=596 · num_kv_indices=546
NVIDIA B200
117.1µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=5422 · num_kv_indices=5372
NVIDIA B200
117.4µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=4334 · num_kv_indices=4284
NVIDIA B200
119.1µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=6510 · num_kv_indices=6460
NVIDIA B200
127.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=7598 · num_kv_indices=7548
NVIDIA B200
129.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=8686 · num_kv_indices=8636
NVIDIA B200
144.8µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=24732 · num_kv_indices=15463
NVIDIA B200
373.3µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=28831 · num_kv_indices=28815
NVIDIA B200
374.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=31007 · num_kv_indices=30991
NVIDIA B200
374.4µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=26908 · num_kv_indices=17639
NVIDIA B200
386.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=33183 · num_kv_indices=33167
NVIDIA B200
386.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=25820 · num_kv_indices=16551
NVIDIA B200
386.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=35359 · num_kv_indices=35343
NVIDIA B200
387.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=39711 · num_kv_indices=39695
NVIDIA B200
389.7µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=27996 · num_kv_indices=18727
NVIDIA B200
398.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=41887 · num_kv_indices=41871
NVIDIA B200
400.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=44063 · num_kv_indices=44047
NVIDIA B200
402.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=46303 · num_kv_indices=46287
NVIDIA B200
402.7µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=48479 · num_kv_indices=48463
NVIDIA B200
404.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=29084 · num_kv_indices=19815
NVIDIA B200
404.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=30172 · num_kv_indices=20903
NVIDIA B200
411.8µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=52831 · num_kv_indices=52815
NVIDIA B200
414.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=50655 · num_kv_indices=50639
NVIDIA B200
416.1µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=55007 · num_kv_indices=54991
NVIDIA B200
417.2µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=31260 · num_kv_indices=21991
NVIDIA B200
425.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=61535 · num_kv_indices=61519
NVIDIA B200
426.7µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=32348 · num_kv_indices=23079
NVIDIA B200
427.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=59359 · num_kv_indices=59343
NVIDIA B200
428.3µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=57183 · num_kv_indices=57167
NVIDIA B200
430.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=37535 · num_kv_indices=37519
NVIDIA B200
430.1µ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:6ff3d45887bd6fffc9365b005926d1286c2456c122012f6932ad59d05a47a9eb
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.
mma
s_ij_subset = tl.dot(q_subset, tl.trans(k.to(tl.float32)))num-warps = 4
num_warps = 4online-softmax
m_new = tl.maximum(m_i, m_ij)stages = 3
num_stages = 3tile-n = 128
BLOCK_N = 128Kernel source
main.py268 lines
import torch
import triton
import triton.language as tl
import math
from typing import Optional
# Constant for converting natural log to base-2 log.
INV_LOG_2 = 1.0 / math.log(2.0)
@triton.jit
def gqa_paged_decode_h32_kv4_d128_ps1_kernel(
# Pointers to Tensors
q_ptr, k_cache_ptr, v_cache_ptr,
kv_indptr_ptr, kv_indices_ptr,
sm_scale,
lse_ptr, output_ptr,
# Stride information
stride_q_bs, stride_q_h,
stride_k_page, stride_k_head,
stride_v_page, stride_v_head,
stride_out_bs, stride_out_h,
stride_lse_bs, stride_lse_h,
# GQA parameters
gqa_ratio: tl.constexpr,
# Meta-parameters
NUM_QO_HEADS: tl.constexpr,
NUM_KV_HEADS: tl.constexpr,
HEAD_DIM: tl.constexpr,
BLOCK_N: tl.constexpr,
INV_LOG_2: tl.constexpr,
):
"""
Triton kernel for paged GQA decode.
Each program computes the attention output for ALL query heads of one sequence.
This is done to satisfy the M >= 16 constraint of tl.dot on B200/H100.
"""
# 1. Get program ID for the batch dimension
pid_b = tl.program_id(0)
# 2. Load sequence information from kv_indptr
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
# 3. Define offsets for head and dimension axes
offs_h = tl.arange(0, NUM_QO_HEADS)
offs_d = tl.arange(0, HEAD_DIM)
# 4. Early exit for sequences with no KV cache
if seq_len == 0:
# Store zero output
output_offset = pid_b * stride_out_bs + offs_h[:, None] * stride_out_h + offs_d[None, :]
tl.store(output_ptr + output_offset, tl.zeros([NUM_QO_HEADS, HEAD_DIM], dtype=tl.bfloat16))
# Store -inf LSE
lse_offset = pid_b * stride_lse_bs + offs_h * stride_lse_h
tl.store(lse_ptr + lse_offset, tl.full([NUM_QO_HEADS], -float('inf'), dtype=tl.float32))
return
# 5. Load Q matrix for all heads of the current sequence
q_offset = pid_b * stride_q_bs + offs_h[:, None] * stride_q_h + offs_d[None, :]
q = tl.load(q_ptr + q_offset).to(tl.float32)
# 6. Initialize accumulators for online softmax (one per head)
# Shapes are [NUM_QO_HEADS, 1] for broadcasting with scores [NUM_QO_HEADS, BLOCK_N]
acc_o = tl.zeros([NUM_QO_HEADS, HEAD_DIM], dtype=tl.float32)
m_i = tl.full([NUM_QO_HEADS, 1], -float('inf'), dtype=tl.float32)
l_i = tl.zeros([NUM_QO_HEADS, 1], dtype=tl.float32)
# 7. Determine the corresponding KV head index for each Q head
kv_head_indices = offs_h // gqa_ratio
# 8. Main loop over the KV sequence length in blocks of BLOCK_N
offs_n = tl.arange(0, BLOCK_N)
kv_indices_base_ptr = kv_indices_ptr + page_start
num_blocks = tl.cdiv(seq_len, BLOCK_N)
for block_idx in range(num_blocks):
# a. Compute offsets and masks for the current block
current_block_start = block_idx * BLOCK_N
kv_indices_offs = current_block_start + offs_n
kv_mask = kv_indices_offs < seq_len
page_ids = tl.load(kv_indices_base_ptr + kv_indices_offs, mask=kv_mask, other=0)
# b. Compute scores S = Q @ K.T
# We iterate over each KV head, compute scores for the corresponding Q heads,
# and accumulate the results.
s_ij = tl.zeros([NUM_QO_HEADS, BLOCK_N], dtype=tl.float32)
for kv_h_idx in range(NUM_KV_HEADS):
# Mask to select Q heads corresponding to the current KV head
q_mask = (kv_head_indices == kv_h_idx)
q_subset = tl.where(q_mask[:, None], q, 0.0)
# Gather load K block for the current KV head
k_ptr = k_cache_ptr + page_ids[:, None] * stride_k_page + \
kv_h_idx * stride_k_head + offs_d[None, :]
k = tl.load(k_ptr, mask=kv_mask[:, None], other=0.0)
# Compute scores for this subset of heads
s_ij_subset = tl.dot(q_subset, tl.trans(k.to(tl.float32)))
s_ij += s_ij_subset
s_ij *= sm_scale
s_ij = tl.where(kv_mask[None, :], s_ij, -float('inf'))
# c. Update online softmax statistics (m_i, l_i)
m_ij = tl.max(s_ij, 1)[:, None]
m_new = tl.maximum(m_i, m_ij)
alpha = tl.exp(m_i - m_new)
beta = tl.exp(s_ij - m_new)
l_i_update = tl.sum(beta, 1)[:, None]
l_i = l_i * alpha + l_i_update
# d. Compute P = softmax(S) and update output accumulator (acc_o)
p_ij = beta.to(tl.bfloat16)
# Rescale old accumulator
acc_o = acc_o * alpha
# Iterate over KV heads again to compute P @ V
for kv_h_idx in range(NUM_KV_HEADS):
# Mask to select probabilities for the current KV head
q_mask = (kv_head_indices == kv_h_idx)
p_ij_subset = tl.where(q_mask[:, None], p_ij, 0.0)
# Gather load V block for the current KV head
v_ptr = v_cache_ptr + page_ids[:, None] * stride_v_page + \
kv_h_idx * stride_v_head + offs_d[None, :]
v = tl.load(v_ptr, mask=kv_mask[:, None], other=0.0)
# Accumulate P @ V for this subset of heads
acc_o += tl.dot(p_ij_subset, v, out_dtype=tl.float32)
m_i = m_new
# 9. Finalize and store output and LSE
# Rescale accumulator to get the final output vector
l_i_safe = tl.where(l_i == 0.0, 1.0, l_i)
o = acc_o / l_i_safe
output_offset = pid_b * stride_out_bs + offs_h[:, None] * stride_out_h + offs_d[None, :]
tl.store(output_ptr + output_offset, o.to(tl.bfloat16))
# Compute and store 2-based log-sum-exp
# The indexing `[:, 0]` is not supported by the Triton compiler in this context.
# Use tl.ravel to flatten the [NUM_QO_HEADS, 1] tensor to [NUM_QO_HEADS]
# before storing, which matches the 1D shape of the destination offsets.
log_lse = m_i + tl.log(l_i)
lse = tl.ravel(log_lse * INV_LOG_2)
lse_offset = pid_b * stride_lse_bs + offs_h * stride_lse_h
tl.store(lse_ptr + lse_offset, lse)
def run(
q: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
kv_indptr: torch.Tensor,
kv_indices: torch.Tensor,
sm_scale: Optional[float] = None,
) -> (torch.Tensor, torch.Tensor):
"""
Wrapper function for the GQA Paged Decode kernel.
Handles device management, tensor validation, kernel launching, and
returning results to the original device.
Args:
q: Query tensor of shape [batch_size, num_qo_heads, head_dim].
k_cache: Key cache tensor of shape [num_pages, page_size, num_kv_heads, head_dim].
v_cache: Value cache tensor of shape [num_pages, page_size, num_kv_heads, head_dim].
kv_indptr: KV page offsets for each sequence, shape [batch_size + 1].
kv_indices: Page IDs for KV cache lookups, shape [num_kv_indices].
sm_scale: Softmax scale factor. Defaults to 1/sqrt(head_dim).
Returns:
A tuple containing:
- output: The attention output tensor of shape [batch_size, num_qo_heads, head_dim].
- lse: The log-sum-exp of attention logits (base 2), shape [batch_size, num_qo_heads].
"""
# 1. --- Device Management & Validation ---
if not torch.cuda.is_available():
raise RuntimeError("Triton kernel requires a CUDA-enabled device.")
original_device = q.device
is_cpu_run = original_device.type == 'cpu'
if is_cpu_run:
# Move all tensor inputs to the default CUDA device
device = "cuda"
q = q.to(device)
k_cache = k_cache.to(device)
v_cache = v_cache.to(device)
kv_indptr = kv_indptr.to(device)
kv_indices = kv_indices.to(device)
else:
# Ensure all tensors are on the same CUDA device
device = q.device
for t_name, t in [("k_cache", k_cache), ("v_cache", v_cache), ("kv_indptr", kv_indptr), ("kv_indices", kv_indices)]:
if t.device != device:
raise ValueError(f"All input tensors must be on the same device. "
f"Expected {device}, but found '{t_name}' on {t.device}.")
# 2. --- Shape and Parameter Validation ---
batch_size, num_qo_heads, head_dim = q.shape
num_pages, page_size, num_kv_heads, _ = k_cache.shape
# Constants from spec
if num_qo_heads != 32: raise ValueError(f"Expected num_qo_heads=32, got {num_qo_heads}")
if num_kv_heads != 4: raise ValueError(f"Expected num_kv_heads=4, got {num_kv_heads}")
if head_dim != 128: raise ValueError(f"Expected head_dim=128, got {head_dim}")
if page_size != 1: raise ValueError(f"Expected page_size=1, got {page_size}")
# Constraints from spec
if kv_indptr.shape != (batch_size + 1,):
raise ValueError(f"Expected kv_indptr shape {(batch_size + 1,)}, got {kv_indptr.shape}")
if kv_indices.shape[0] != kv_indptr[-1].item():
raise ValueError(f"Mismatch in total number of KV indices.")
if sm_scale is None:
sm_scale = 1.0 / math.sqrt(head_dim)
# 3. --- Kernel Launch Setup ---
output = torch.empty_like(q)
lse = torch.empty((batch_size, num_qo_heads), dtype=torch.float32, device=device)
gqa_ratio = num_qo_heads // num_kv_heads
# Each program handles all heads for one batch item
grid = (batch_size,)
# Kernel meta-parameters optimized for B200
BLOCK_N = 128
num_warps = 4
num_stages = 3
# 4. --- Launch Kernel ---
gqa_paged_decode_h32_kv4_d128_ps1_kernel[grid](
q, k_cache, v_cache, kv_indptr, kv_indices,
float(sm_scale),
lse, output,
# Strides
q.stride(0), q.stride(1),
k_cache.stride(0), k_cache.stride(2),
v_cache.stride(0), v_cache.stride(2),
output.stride(0), output.stride(1),
lse.stride(0), lse.stride(1),
# GQA parameters
gqa_ratio=gqa_ratio,
# Meta-parameters
NUM_QO_HEADS=num_qo_heads,
NUM_KV_HEADS=num_kv_heads,
HEAD_DIM=head_dim,
BLOCK_N=BLOCK_N,
INV_LOG_2=INV_LOG_2,
num_warps=num_warps,
num_stages=num_stages,
)
# 5. --- Return Results ---
if is_cpu_run:
# Move results back to the original CPU device
output = output.to(original_device)
lse = lse.to(original_device)
return output, lsescrolls · 268 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON