gemini-2.5-pro / triton3nob6q
gemini-2.5-pro_triton_3nob6q · gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 255 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-3nob6q?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:3f88e26faec95e1c346dfe300a637a435a465b15bceecb537f747397fc559e6c
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.
num-warps = 8
num_warps = 8online-softmax
m_new = tl.maximum(m, tl.max(s, axis=0))tile-n = 128
BLOCK_N = 128Kernel source
main.py255 lines
import torch
import triton
import triton.language as tl
import math
# Wrapper for device management
def _get_device_and_wrapper(args, kwargs):
"""
Finds a common device for all tensors and returns a wrapper function
to move results back to the original device.
"""
device = None
original_devices = {}
def find_device(tensor, name):
nonlocal device
if isinstance(tensor, torch.Tensor):
if name not in original_devices:
original_devices[name] = tensor.device
if device is None:
device = tensor.device
elif tensor.device != device:
raise ValueError(f"All tensors must be on the same device. Expected {device}, but got {tensor.device} for {name}.")
# Process args and kwargs to find the target device
for i, arg in enumerate(args):
find_device(arg, f"arg_{i}")
for k, v in kwargs.items():
find_device(v, k)
if device is None:
# No tensors found, default to CUDA if available, else CPU
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
if device.type == "cpu" and torch.cuda.is_available():
# Move CPU tensors to GPU if CUDA is available
target_device = torch.device("cuda")
else:
target_device = device
if target_device.type != 'cuda':
raise RuntimeError("Triton kernels require a CUDA-enabled GPU.")
def to_device(o, name):
if isinstance(o, torch.Tensor) and o.device != target_device:
return o.to(target_device)
return o
processed_args = [to_device(arg, f"arg_{i}") for i, arg in enumerate(args)]
processed_kwargs = {k: to_device(v, k) for k, v in kwargs.items()}
def unwrap(result):
if isinstance(result, torch.Tensor):
# Restore to the device of the first tensor input 'q'.
original_dev = original_devices.get('q', device)
return result.to(original_dev)
elif isinstance(result, (list, tuple)):
return type(result)(unwrap(item) for item in result)
return result
return target_device, processed_args, processed_kwargs, unwrap
@triton.jit
def _kernel(
# Inputs
Q, K_cache, V_cache,
qo_indptr, kv_indptr, kv_indices,
q_to_b_map,
sm_scale,
# Outputs
O, LSE,
# Strides
stride_q_token, stride_q_head,
stride_k_page, stride_k_head,
stride_v_page, stride_v_head,
stride_o_token, stride_o_head,
stride_lse_token,
# Constants
N_Q_HEADS: tl.constexpr,
N_KV_HEADS: tl.constexpr,
HEAD_DIM: tl.constexpr,
PAGE_SIZE: tl.constexpr,
GQA_RATIO: tl.constexpr,
BLOCK_D: tl.constexpr,
BLOCK_N: tl.constexpr,
):
"""
Triton kernel for GQA paged prefill with causal masking.
Each program instance computes attention for one query token and one query head.
"""
# Grid: each program handles one query token and one query head
q_token_idx = tl.program_id(0)
q_head_idx = tl.program_id(1)
# 1. Look up batch index and sequence properties using the precomputed map
b_idx = tl.load(q_to_b_map + q_token_idx)
q_start = tl.load(qo_indptr + b_idx)
kv_start = tl.load(kv_indptr + b_idx)
kv_end = tl.load(kv_indptr + b_idx + 1)
num_kv_tokens = kv_end - kv_start
# 2. Calculate causal mask limit for the current query token
q_seq_offset = q_token_idx - q_start
num_q_tokens = tl.load(qo_indptr + b_idx + 1) - q_start
delta = num_kv_tokens - num_q_tokens
causal_limit = q_seq_offset + 1 + delta
# The max number of KV tokens to attend to is limited by both causality
# and the actual number of KV tokens available in the sequence.
max_kv_len = tl.minimum(causal_limit, num_kv_tokens)
# Ensure max_kv_len is not negative, which can happen if causal_limit is negative.
max_kv_len = tl.maximum(0, max_kv_len)
# 3. Load query vector
d_offs = tl.arange(0, BLOCK_D)
q_ptr = Q + q_token_idx * stride_q_token + q_head_idx * stride_q_head
q = tl.load(q_ptr + d_offs, mask=d_offs < HEAD_DIM, other=0.0).to(tl.float32)
# 4. Initialize accumulators for online softmax
m = -float("inf")
l = 0.0
acc = tl.zeros([BLOCK_D], dtype=tl.float32)
# 5. Determine corresponding KV head for GQA
kv_head_idx = q_head_idx // GQA_RATIO
# 6. Loop over KV sequence in blocks of size BLOCK_N
kv_indices_base_ptr = kv_indices + kv_start
k_block_start = 0
while k_block_start < max_kv_len:
kv_seq_offs = k_block_start + tl.arange(0, BLOCK_N)
kv_mask = kv_seq_offs < max_kv_len
# Load page IDs for the current block from kv_indices
page_ids = tl.load(kv_indices_base_ptr + kv_seq_offs, mask=kv_mask, other=0)
# Construct pointers for indirect access to K and V caches
d_offs_exp = d_offs[None, :]
k_ptrs = K_cache + (page_ids[:, None] * stride_k_page + kv_head_idx * stride_k_head + d_offs_exp)
v_ptrs = V_cache + (page_ids[:, None] * stride_v_page + kv_head_idx * stride_v_head + d_offs_exp)
# Load K and V blocks
block_mask = kv_mask[:, None] & (d_offs[None, :] < HEAD_DIM)
k = tl.load(k_ptrs, mask=block_mask, other=0.0)
v = tl.load(v_ptrs, mask=block_mask, other=0.0)
# --- Compute attention scores (S = Q @ K.T) ---
s = tl.sum(q[None, :] * k.to(tl.float32), axis=1) * sm_scale
s = tl.where(kv_mask, s, -float("inf"))
# --- Online softmax update ---
m_new = tl.maximum(m, tl.max(s, axis=0))
p = tl.exp(s - m_new)
l_new = tl.exp(m - m_new) * l + tl.sum(p, axis=0)
# --- Update accumulator (acc) ---
acc = acc * tl.exp(m - m_new)
p = p.to(v.dtype)
acc += tl.sum(p[:, None] * v, axis=0)
# Update state and advance to the next block
m = m_new
l = l_new
k_block_start += BLOCK_N
# 7. Finalize output and LSE
o = acc / tl.where(l == 0.0, 1.0, l)
# If l is 0, m is -inf, and log(l) is -inf. Result is correctly -inf.
lse = m + tl.log(l)
# Convert to 2-based log-sum-exp as per spec
LOG2E = 1.4426950408889634
lse *= LOG2E
# 8. Write results to global memory
o_ptr = O + q_token_idx * stride_o_token + q_head_idx * stride_o_head
lse_ptr = LSE + q_token_idx * stride_lse_token + q_head_idx
tl.store(o_ptr + d_offs, o.to(O.dtype.element_ty), mask=d_offs < HEAD_DIM)
tl.store(lse_ptr, lse)
def gqa_paged_prefill_causal_h32_kv4_d128_ps1(q, k_cache, v_cache, qo_indptr, kv_indptr, kv_indices, sm_scale):
"""
Computes Grouped-Query Attention for a batch of sequences with paged KV cache
and causal masking, optimized for prefill phase.
"""
# 1. Extract dimensions and constants from inputs
total_q, num_qo_heads, head_dim = q.shape
num_pages, page_size, num_kv_heads, _ = k_cache.shape
# 2. Assertions to ensure shapes match the spec
assert num_qo_heads == 32 and num_kv_heads == 4
assert head_dim == 128 and page_size == 1
assert total_q == qo_indptr[-1].item()
assert kv_indptr[-1].item() == kv_indices.shape[0]
# 3. Pre-computation on host: map each query token to its batch index
q_starts = qo_indptr[:-1]
seq_lens = qo_indptr[1:] - q_starts
batch_size = len(seq_lens)
b_indices = torch.arange(batch_size, device=q.device, dtype=torch.int32)
q_to_b_map = torch.repeat_interleave(b_indices, seq_lens.to(torch.long))
# 4. Allocate output tensors
output = torch.empty_like(q)
lse = torch.empty((total_q, num_qo_heads), dtype=torch.float32, device=q.device)
# 5. Set up Triton grid
grid = (total_q, num_qo_heads)
# 6. Define constants for the kernel
GQA_RATIO = num_qo_heads // num_kv_heads
BLOCK_N = 128
num_warps = 8
# 7. Launch the Triton kernel
_kernel[grid](
q, k_cache, v_cache,
qo_indptr, kv_indptr, kv_indices,
q_to_b_map,
sm_scale,
output, lse,
q.stride(0), q.stride(1),
k_cache.stride(0), k_cache.stride(2), # Stride over num_kv_heads
v_cache.stride(0), v_cache.stride(2), # Stride over num_kv_heads
output.stride(0), output.stride(1),
lse.stride(0),
N_Q_HEADS=num_qo_heads,
N_KV_HEADS=num_kv_heads,
HEAD_DIM=head_dim,
PAGE_SIZE=page_size,
GQA_RATIO=GQA_RATIO,
BLOCK_D=head_dim,
BLOCK_N=BLOCK_N,
num_warps=num_warps
)
return output, lse
def run(*args, **kwargs):
"""
Public entry point for the kernel.
Handles device management and calls the main implementation.
"""
target_device, processed_args, processed_kwargs, unwrap_fn = _get_device_and_wrapper(args, kwargs)
result = gqa_paged_prefill_causal_h32_kv4_d128_ps1(*processed_args, **processed_kwargs)
return unwrap_fn(result)scrolls · 255 lines total
Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0
Best evidence level for this revision: reported
JSON