gemini-2.5-pro / tritonrbz3hy
gemini-2.5-pro_triton_rbz3hy · gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 300 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-rbz3hy?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32
Benchmark evidence
No published measurement for this revision.
No evidence · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:dbbee34fd020694337a0db4c42071daf911f49664f91fdb56a170471ade15444
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(mma
s = tl.dot(q_mat, tl.trans(k)) * sm_scalenum-warps = 4
triton.Config({'BLOCK_N': 64}, num_warps=4, num_stages=3),online-softmax
m_i_new = tl.maximum(m_i, tl.max(s, axis=1))stages = 3
triton.Config({'BLOCK_N': 64}, num_warps=4, num_stages=3),Kernel source
main.py300 lines
import torch
import triton
import triton.language as tl
import math
# `triton.jit`'ed functions can be auto-tuned by using the `triton.autotune` decorator, which consumes:
# - A list of `triton.Config` objects that define different configurations of values for user-defined arguments.
# - A `key` argument containing a list of names of arguments used to determine which configuration is chosen.
@triton.autotune(
configs=[
triton.Config({'BLOCK_N': 64}, num_warps=4, num_stages=3),
triton.Config({'BLOCK_N': 128}, num_warps=4, num_stages=3),
triton.Config({'BLOCK_N': 256}, num_warps=8, num_stages=2),
triton.Config({'BLOCK_N': 128}, num_warps=8, num_stages=4),
triton.Config({'BLOCK_N': 64}, num_warps=4, num_stages=4),
triton.Config({'BLOCK_N': 32}, num_warps=2, num_stages=2),
],
key=['HEAD_DIM'],
)
@triton.jit
def gqa_ragged_prefill_causal_kernel(
# Pointers to matrices
Q, K, V, O, LSE,
# Pointer to precomputed location map
q_loc,
sm_scale,
# Strides
Q_stride_t, Q_stride_h, Q_stride_d,
K_stride_t, K_stride_h, K_stride_d,
V_stride_t, V_stride_h, V_stride_d,
O_stride_t, O_stride_h, O_stride_d,
LSE_stride_t, LSE_stride_h,
q_loc_stride_t, q_loc_stride_d,
# Compile-time constants
NUM_QO_HEADS: tl.constexpr,
NUM_KV_HEADS: tl.constexpr,
HEAD_DIM: tl.constexpr,
BLOCK_D: tl.constexpr,
BLOCK_N: tl.constexpr,
):
"""
Triton kernel for Grouped-Query Attention for ragged prefill with causal masking.
This kernel computes attention for one query token against its corresponding
key-value sequence. The grid is launched with one program per (query_token, query_head).
The ragged nature of the input is handled by a precomputed location map `q_loc`,
which provides sequence boundaries for each query token, avoiding complex and slow
indexing logic inside the kernel.
The computation uses a tiled approach similar to FlashAttention to efficiently
process the key-value sequence in blocks, leveraging shared memory implicitly
via Triton's dot product operations. Online softmax is used to maintain
numerical stability and compute the result in a single pass over the KV cache.
"""
# Grid is (total_q, num_qo_heads)
pid_qt = tl.program_id(0) # Global query token index
pid_h = tl.program_id(1) # Query head index
# --- 1. Load sequence boundaries and determine context ---
# Load [q_start, q_end, kv_start, kv_end] from the precomputed map
q_loc_ptr = q_loc + pid_qt * q_loc_stride_t
q_start = tl.load(q_loc_ptr + 0 * q_loc_stride_d)
q_end = tl.load(q_loc_ptr + 1 * q_loc_stride_d)
kv_start = tl.load(q_loc_ptr + 2 * q_loc_stride_d)
kv_end = tl.load(q_loc_ptr + 3 * q_loc_stride_d)
# Calculate local query index and sequence lengths
q_idx_local = pid_qt - q_start
num_q_tokens = q_end - q_start
num_kv_tokens = kv_end - kv_start
delta = num_kv_tokens - num_q_tokens
# Causal sequence length for this query
kv_len_for_q = tl.minimum(q_idx_local + 1 + delta, num_kv_tokens)
# --- 2. Determine head indices and pointers ---
GQA_RATIO: tl.constexpr = NUM_QO_HEADS // NUM_KV_HEADS
kv_head_idx = pid_h // GQA_RATIO
# Pointers to K and V for the correct sequence and head
k_batch_head_ptr = K + kv_start * K_stride_t + kv_head_idx * K_stride_h
v_batch_head_ptr = V + kv_start * V_stride_t + kv_head_idx * V_stride_h
# --- 3. Initialize accumulator and online softmax statistics ---
acc = tl.zeros([BLOCK_D], dtype=tl.float32)
m_i = -float('inf')
l_i = 0.0
# --- 4. Load query vector ---
q_ptr = Q + pid_qt * Q_stride_t + pid_h * Q_stride_h
offs_d = tl.arange(0, BLOCK_D)
q = tl.load(q_ptr + offs_d, mask=offs_d < HEAD_DIM, other=0.0).to(tl.float32)
# --- 5. Main loop over KV sequence blocks ---
kv_offset = 0
# The loop condition handles cases where kv_len_for_q <= 0
while kv_offset < kv_len_for_q:
# Pointers to the current block of K and V
k_ptr = k_batch_head_ptr + kv_offset * K_stride_t
v_ptr = v_batch_head_ptr + kv_offset * V_stride_t
# Offsets for loading K and V blocks
offs_n = tl.arange(0, BLOCK_N)
k_offs = (offs_n[:, None] * K_stride_t + offs_d[None, :])
v_offs = (offs_n[:, None] * V_stride_t + offs_d[None, :])
# Create a mask for the current block to handle both padding within the
# block and the causal boundary.
k_mask = (kv_offset + offs_n) < kv_len_for_q
# Load K and V blocks with masking
k = tl.load(k_ptr + k_offs, mask=k_mask[:, None] & (offs_d[None, :] < HEAD_DIM), other=0.0).to(tl.float32)
v = tl.load(v_ptr + v_offs, mask=k_mask[:, None] & (offs_d[None, :] < HEAD_DIM), other=0.0)
# --- Core attention computation (FIXED BLOCK) ---
# Compute Q @ K.T
# FIX: The core issue was that tl.dot requires 2D inputs. We reshape the
# 1D query vector `q` into a 2D matrix `q_mat` of shape [1, BLOCK_D].
q_mat = tl.reshape(q, (1, BLOCK_D))
s = tl.dot(q_mat, tl.trans(k)) * sm_scale
# Apply mask to logits. s is [1, N], k_mask is [N], broadcasting is fine.
s = tl.where(k_mask, s, -float('inf'))
# --- Online softmax update ---
# FIX: Since `s` is now 2D [1, N], the reductions must be handled correctly.
# We reduce over axis=1 to get [1]-shaped tensors for the statistics.
m_i_new = tl.maximum(m_i, tl.max(s, axis=1))
alpha = tl.exp(m_i - m_i_new)
p = tl.exp(s - m_i_new)
l_i_new = alpha * l_i + tl.sum(p, axis=1)
# Update accumulator
# Triton correctly broadcasts the [1]-shaped `alpha` tensor across `acc`.
acc = acc * alpha
# FIX: `p` is [1, N], `v` is [N, D]. Dot product gives [1, D].
# We must reshape the result to [D] to correctly add it to `acc`.
delta_acc = tl.dot(p.to(v.dtype), v)
acc += tl.reshape(delta_acc, (BLOCK_D,))
# Update statistics for next iteration
m_i = m_i_new
l_i = l_i_new
kv_offset += BLOCK_N
# --- 6. Finalize and store results ---
# Finalize accumulator
l_i_safe = tl.where(l_i == 0.0, 1.0, l_i)
acc = acc / l_i_safe
# Compute 2-based log-sum-exp
log2_e = 1.4426950408889634 # 1.0 / ln(2)
lse = m_i + tl.log(l_i)
lse = lse * log2_e
# If all scores were -inf, l_i is 0, log(l_i) is -inf, which is correct.
lse = tl.where(l_i == 0.0, -float('inf'), lse)
# Store output and LSE
offs_d_store = tl.arange(0, BLOCK_D)
o_ptr = O + pid_qt * O_stride_t + pid_h * O_stride_h
lse_ptr = LSE + pid_qt * LSE_stride_t + pid_h * LSE_stride_h
tl.store(o_ptr + offs_d_store, acc.to(O.dtype.element_ty), mask=offs_d_store < HEAD_DIM)
tl.store(lse_ptr, lse)
def gqa_ragged_prefill_causal_h32_kv8_d128(q, k, v, qo_indptr, kv_indptr, sm_scale):
"""
Wrapper function for the GQA ragged prefill kernel.
This function prepares tensors, defines the launch grid, and calls the
Triton kernel. It includes a host-side precomputation step to create
a `q_loc` map, which massively simplifies the kernel's indexing logic
by providing each query token with its sequence boundaries directly.
"""
# Extract shape information
total_q, num_qo_heads, head_dim = q.shape
total_kv, num_kv_heads, _ = k.shape
len_indptr = qo_indptr.shape[0]
batch_size = len_indptr - 1
# Allocate output tensors
output = torch.empty_like(q)
lse = torch.empty((total_q, num_qo_heads), dtype=torch.float32, device=q.device)
# Precompute location map: [q_start, q_end, kv_start, kv_end] for each query
# This avoids a complex and slow search/lookup within the kernel.
q_loc = torch.empty((total_q, 4), dtype=torch.int32, device=q.device)
if batch_size > 0 and total_q > 0:
for b in range(batch_size):
q_s, q_e = qo_indptr[b].item(), qo_indptr[b+1].item()
kv_s, kv_e = kv_indptr[b].item(), kv_indptr[b+1].item()
if q_s < q_e:
# Use broadcasting to fill the map for all tokens in the sequence
q_loc[q_s:q_e, 0] = q_s
q_loc[q_s:q_e, 1] = q_e
q_loc[q_s:q_e, 2] = kv_s
q_loc[q_s:q_e, 3] = kv_e
# Define the launch grid: one program per (query_token, query_head)
grid = (total_q, num_qo_heads)
# Call the Triton kernel, only if there are tokens to process
if total_q > 0:
gqa_ragged_prefill_causal_kernel[grid](
q, k, v, output, lse,
q_loc,
sm_scale,
# Strides
q.stride(0), q.stride(1), q.stride(2),
k.stride(0), k.stride(1), k.stride(2),
v.stride(0), v.stride(1), v.stride(2),
output.stride(0), output.stride(1), output.stride(2),
lse.stride(0), lse.stride(1),
q_loc.stride(0), q_loc.stride(1),
# Constants
NUM_QO_HEADS=num_qo_heads,
NUM_KV_HEADS=num_kv_heads,
HEAD_DIM=head_dim,
BLOCK_D=head_dim,
# BLOCK_N is autotuned
)
return output, lse
def run(*args, **kwargs):
"""
Public entry point for the operation.
Handles device management, argument parsing, and calls the underlying
Triton implementation. It ensures that input tensors are on the correct
device (CUDA) and that output tensors are moved back to the original
device.
"""
# --- Argument Parsing ---
# This robustly handles both positional and keyword arguments.
arg_names = ['q', 'k', 'v', 'qo_indptr', 'kv_indptr', 'sm_scale']
arg_dict = {name: kwargs.get(name) for name in arg_names}
for i, arg in enumerate(args):
# This will overwrite a kwarg if it was also passed as an arg,
# which is standard Python behavior.
if i < len(arg_names):
arg_dict[arg_names[i]] = arg
q = arg_dict['q']
k = arg_dict['k']
v = arg_dict['v']
qo_indptr = arg_dict['qo_indptr']
kv_indptr = arg_dict['kv_indptr']
sm_scale = arg_dict['sm_scale']
# Check for missing required arguments
required_args = ['q', 'k', 'v', 'qo_indptr', 'kv_indptr']
for arg_name in required_args:
if arg_dict[arg_name] is None:
raise TypeError(f"Missing required argument: '{arg_name}'")
# --- Constants and Defaults ---
HEAD_DIM = 128
if sm_scale is None:
sm_scale = 1.0 / math.sqrt(HEAD_DIM)
# --- Device Management ---
if not torch.cuda.is_available():
raise RuntimeError("This Triton kernel requires a CUDA-enabled GPU.")
original_device = q.device
is_cpu = original_device.type == 'cpu'
if is_cpu:
# Move all tensors to GPU
q, k, v, qo_indptr, kv_indptr = (
t.cuda() for t in [q, k, v, qo_indptr, kv_indptr]
)
elif q.device.type != 'cuda':
raise RuntimeError(f"Unsupported device: {q.device}. Only CPU and CUDA are supported.")
# --- Constraints Validation ---
total_q, num_qo_heads, head_dim = q.shape
_, num_kv_heads, _ = k.shape
assert num_qo_heads == 32, f"Expected num_qo_heads=32, but got {num_qo_heads}"
assert num_kv_heads == 8, f"Expected num_kv_heads=8, but got {num_kv_heads}"
assert head_dim == HEAD_DIM, f"Expected head_dim={HEAD_DIM}, but got {head_dim}"
assert total_q == qo_indptr[-1].item(), "total_q must match qo_indptr[-1]"
assert k.shape[0] == kv_indptr[-1].item(), "total_kv must match kv_indptr[-1]"
# --- Kernel Execution ---
output, lse = gqa_ragged_prefill_causal_h32_kv8_d128(q, k, v, qo_indptr, kv_indptr, sm_scale)
# --- Device Restoration ---
if is_cpu:
output = output.to(original_device)
lse = lse.to(original_device)
return output, lsescrolls · 300 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON