gemini-2.5-pro_triton_zezbpc
gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 239 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-zezbpc?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32
Benchmark evidence
15 measurements across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
GQA ragged prefill causal h32 kv4 d128bf16 · [1, 4, 128] · #7c206f
NVIDIA B200
97.6µs
#4 of 5
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [6, 4, 128] · #6d6644
NVIDIA B200
97.8µs
#4 of 5
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [34, 4, 128] · #816a2c
NVIDIA B200
100.4µs
#2 of 5
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [34, 4, 128] · #641e77
NVIDIA B200
103.9µs
#5 of 20
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [6, 4, 128] · #55a16d
NVIDIA B200
104.1µs
#7 of 10
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [6, 4, 128] · #55a16d
NVIDIA B200
104.6µs
#8 of 10
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [34, 4, 128] · #641e77
NVIDIA B200
104.7µs
#6 of 20
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [1, 4, 128] · #ce8167
NVIDIA B200
105.2µs
#7 of 10
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [1, 4, 128] · #ce8167
NVIDIA B200
105.4µs
#8 of 10
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [34, 4, 128] · #641e77
NVIDIA B200
106.1µs
#7 of 20
2025-10-20
Show all 15 measurements ›Showing all 15 measurements ⌄
GQA ragged prefill causal h32 kv4 d128bf16 · [34, 4, 128] · #641e77
NVIDIA B200
106.3µs
#8 of 20
2025-10-20
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:6c278eec14ff6845c07c4715a9cbb85d5bcaf49d7b90ecfbd3baf5b1787a1641
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.
tile-n = 64
BLOCK_N = 64Kernel source
main.py239 lines
import torch
import triton
import triton.language as tl
import math
@triton.jit
def gqa_ragged_prefill_causal_h32_kv4_d128_kernel(
# Pointers to tensors
q_ptr, k_ptr, v_ptr,
qo_indptr_ptr, kv_indptr_ptr, q_to_b_idx_ptr,
output_ptr, lse_ptr,
# Scalar
sm_scale,
# Strides
q_stride_tq, q_stride_h,
k_stride_tk, k_stride_h,
v_stride_tk, v_stride_h,
# Other metadata
total_q,
# Constants for clarity and performance
GQA_RATIO: tl.constexpr,
NUM_QO_HEADS: tl.constexpr,
# Compile-time constants
HEAD_DIM: tl.constexpr,
BLOCK_N: tl.constexpr,
):
"""
Triton kernel for Grouped-Query Attention on ragged tensors for prefill.
This kernel is specialized for causal attention with specific head dimensions.
Each program instance computes the attention output for one query token and one query head.
"""
# Get program IDs to identify the current query token and head
global_q_idx = tl.program_id(0)
h_qo_idx = tl.program_id(1)
# Find the sequence (batch element) index for the current query token
b_idx = tl.load(q_to_b_idx_ptr + global_q_idx)
# Load sequence boundaries from indptr tensors
q_start = tl.load(qo_indptr_ptr + b_idx)
q_end = tl.load(qo_indptr_ptr + b_idx + 1)
kv_start = tl.load(kv_indptr_ptr + b_idx)
kv_end = tl.load(kv_indptr_ptr + b_idx + 1)
# Calculate causal attention length limit
q_idx_in_seq = global_q_idx - q_start
delta = (kv_end - kv_start) - (q_end - q_start)
max_kv_len = q_idx_in_seq + 1 + delta
# Initialize accumulators for online softmax
m_i = -float('inf')
l_i = 0.0
acc = tl.zeros([HEAD_DIM], dtype=tl.float32)
# Determine the corresponding KV head for the current QO head
h_kv_idx = h_qo_idx // GQA_RATIO
# Load the query vector
d_offsets = tl.arange(0, HEAD_DIM)
q_offset = global_q_idx * q_stride_tq + h_qo_idx * q_stride_h
q_ptrs = q_ptr + q_offset + d_offsets
q_vec = tl.load(q_ptrs).to(tl.float32)
# Loop over the key/value sequence in blocks
num_n_blocks = (max_kv_len + BLOCK_N - 1) // BLOCK_N
for block_n_idx in range(num_n_blocks):
# --- Compute offsets and mask for the current block of K/V ---
kv_idx_in_seq_start = block_n_idx * BLOCK_N
n_offsets = kv_idx_in_seq_start + tl.arange(0, BLOCK_N)
kv_mask = n_offsets < max_kv_len
global_kv_indices = kv_start + n_offsets
# --- Load K block ---
k_offset = global_kv_indices * k_stride_tk + h_kv_idx * k_stride_h
k_ptrs = k_ptr + k_offset[:, None] + d_offsets[None, :]
k_block = tl.load(k_ptrs, mask=kv_mask[:, None], other=0.0).to(tl.float32)
# --- Compute S = Q @ K.T ---
s_block = tl.sum(q_vec[None, :] * k_block, axis=1)
s_block = s_block * sm_scale
s_block = tl.where(kv_mask, s_block, -float('inf'))
# --- Online softmax update ---
m_i_prev = m_i
m_i = tl.maximum(m_i, tl.max(s_block, axis=0))
p = tl.exp(s_block - m_i)
l_i = l_i * tl.exp(m_i_prev - m_i) + tl.sum(p, axis=0)
# --- Load V block and update accumulator ---
v_offset = global_kv_indices * v_stride_tk + h_kv_idx * v_stride_h
v_ptrs = v_ptr + v_offset[:, None] + d_offsets[None, :]
v_block = tl.load(v_ptrs, mask=kv_mask[:, None], other=0.0).to(tl.float32)
# Rescale accumulator before adding new values
acc = acc * tl.exp(m_i_prev - m_i)
# FIX: The original tl.dot(p, v_block) caused a compilation error because `p` is 1D
# while tl.dot requires 2D inputs for matrix multiplication.
# The correct operation is a weighted sum of value vectors: sum(p[i] * v_block[i]).
# This is implemented by reshaping p to [BLOCK_N, 1] for broadcasting,
# multiplying with v_block, and then summing over the block dimension (axis=0).
acc += tl.sum(p[:, None] * v_block, axis=0)
# Finalize and store output vector
# Guard against division by zero if l_i is 0 (e.g., empty sequence)
o = tl.where(l_i > 0, acc / l_i, 0.0)
output_offset = global_q_idx * q_stride_tq + h_qo_idx * q_stride_h
output_ptrs = output_ptr + output_offset + d_offsets
tl.store(output_ptrs, o.to(tl.bfloat16))
# Finalize and store log-sum-exp (LSE)
LOG2_E = 1.4426950408889634 # 1.0 / math.log(2.0)
# Guard against log(0)
lse = m_i + tl.log(l_i + 1e-9)
lse = lse * LOG2_E
lse_offset = global_q_idx * NUM_QO_HEADS + h_qo_idx
tl.store(lse_ptr + lse_offset, lse)
def _get_device(*tensors):
"""
Gets the common device of a list of tensors, handling CPU/CUDA logic.
"""
devices = {t.device.type for t in tensors if hasattr(t, 'device')}
if not devices:
return torch.device('cpu')
if 'cuda' in devices:
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available, but input tensors are on CUDA.")
cuda_devices = {t.device for t in tensors if t.device.type == 'cuda'}
if len(cuda_devices) > 1:
raise RuntimeError(f"Input tensors are on multiple CUDA devices: {cuda_devices}")
return list(cuda_devices)[0]
if torch.cuda.is_available():
return torch.device('cuda')
else:
raise RuntimeError("Triton kernels require a CUDA-enabled GPU, but none was found.")
def run(*args, **kwargs):
"""
Entry point for the GQA Ragged Prefill Causal Attention kernel.
Args:
q (torch.Tensor): Query tensor of shape [total_q, num_qo_heads, head_dim].
k (torch.Tensor): Key tensor of shape [total_kv, num_kv_heads, head_dim].
v (torch.Tensor): Value tensor of shape [total_kv, num_kv_heads, head_dim].
qo_indptr (torch.Tensor): Query offsets for each sequence of shape [len_indptr].
kv_indptr (torch.Tensor): Key-value offsets for each sequence of shape [len_indptr].
sm_scale (float, optional): Softmax scale. Defaults to 1/sqrt(head_dim).
Returns:
Tuple[torch.Tensor, torch.Tensor]:
- output (torch.Tensor): Attention output of shape [total_q, num_qo_heads, head_dim].
- lse (torch.Tensor): Log-sum-exp of attention logits of shape [total_q, num_qo_heads].
"""
# 1. Argument parsing
arg_names = ['q', 'k', 'v', 'qo_indptr', 'kv_indptr', 'sm_scale']
expected_arg_count = 5
if len(args) > len(arg_names):
raise TypeError(f"run() takes at most {len(arg_names)} positional arguments but {len(args)} were given")
params = {name: val for name, val in zip(arg_names, args)}
params.update(kwargs)
missing_args = [name for name in arg_names[:expected_arg_count] if name not in params]
if missing_args:
raise TypeError(f"run() missing {len(missing_args)} required positional argument(s): {', '.join(missing_args)}")
q, k, v, qo_indptr, kv_indptr = [params[name] for name in arg_names[:expected_arg_count]]
sm_scale = params.get('sm_scale')
# 2. Constants and shape assertions
NUM_QO_HEADS = 32
NUM_KV_HEADS = 4
HEAD_DIM = 128
total_q, num_qo_heads, head_dim = q.shape
total_kv, num_kv_heads, _ = k.shape
len_indptr = qo_indptr.shape[0]
assert num_qo_heads == NUM_QO_HEADS, f"Expected num_qo_heads={NUM_QO_HEADS}, got {num_qo_heads}"
assert num_kv_heads == NUM_KV_HEADS, f"Expected num_kv_heads={NUM_KV_HEADS}, got {num_kv_heads}"
assert head_dim == HEAD_DIM, f"Expected head_dim={HEAD_DIM}, got {head_dim}"
assert qo_indptr.dim() == 1 and kv_indptr.dim() == 1, "indptr tensors must be 1D"
assert len_indptr > 0, "indptr tensors cannot be empty"
assert total_q == qo_indptr[-1].item(), f"total_q ({total_q}) must match qo_indptr[-1] ({qo_indptr[-1].item()})"
assert total_kv == kv_indptr[-1].item(), f"total_kv ({total_kv}) must match kv_indptr[-1] ({kv_indptr[-1].item()})"
assert qo_indptr.shape == kv_indptr.shape, "qo_indptr and kv_indptr must have the same shape"
if sm_scale is None:
sm_scale = 1.0 / math.sqrt(HEAD_DIM)
# 3. Device management
initial_device = q.device
kernel_device = _get_device(q, k, v, qo_indptr, kv_indptr)
q, k, v, qo_indptr, kv_indptr = [t.to(kernel_device) for t in [q, k, v, qo_indptr, kv_indptr]]
q, k, v = [t.contiguous() for t in [q, k, v]]
# 4. Prepare kernel inputs and outputs
output = torch.empty_like(q, dtype=torch.bfloat16)
lse = torch.full((total_q, NUM_QO_HEADS), -float("inf"), dtype=torch.float32, device=kernel_device)
# 5. Launch kernel
grid = (total_q, NUM_QO_HEADS)
BLOCK_N = 64
if total_q > 0:
# Precompute a mapping from global query index to batch index for efficient lookup in the kernel
q_indices = torch.arange(total_q, device=kernel_device)
qo_ends = qo_indptr[1:]
q_to_b_idx = torch.searchsorted(qo_ends, q_indices, right=True)
gqa_ragged_prefill_causal_h32_kv4_d128_kernel[grid](
q, k, v,
qo_indptr, kv_indptr, q_to_b_idx,
output, lse,
sm_scale,
q.stride(0), q.stride(1),
k.stride(0), k.stride(1),
v.stride(0), v.stride(1),
total_q,
GQA_RATIO=NUM_QO_HEADS // NUM_KV_HEADS,
NUM_QO_HEADS=NUM_QO_HEADS,
HEAD_DIM=HEAD_DIM,
BLOCK_N=BLOCK_N,
)
# 6. Restore output device
output = output.to(initial_device)
lse = lse.to(initial_device)
return output, lsescrolls · 239 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON