gemini-2.5-pro / tritonxvhq2i
gemini-2.5-pro_triton_xvhq2i · gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 318 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-xvhq2i?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:9110256ab67b36605dd4f23f0694760660a8844fd3d086fb12ee676ea9a96d4e
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_CKV': 64, 'BLOCK_KPE': 64, 'BLOCK_KV': 64}, num_stages=3, num_warps=4),online-softmax
m_i_new = tl.maximum(m_i, tl.max(logits, axis=0))stages = 3
triton.Config({'BLOCK_CKV': 64, 'BLOCK_KPE': 64, 'BLOCK_KV': 64}, num_stages=3, num_warps=4),Kernel source
main.py318 lines
import torch
import triton
import triton.language as tl
import math
import inspect
#
# Triton kernel for paged prefill attention
#
# This kernel is optimized for a specific attention variant:
# - Causal attention for prefill (each query attends to keys up to its own position).
# - Paged KV cache with page_size = 1, meaning each entry in kv_indices points to a single token's KV state.
# - Mixed-Logit Attention: Logits are computed from two separate dot products, one for the main content (ckv) and one for positional embeddings (kpe).
# `logits = (q_nope @ K_ckv.T) + (q_pe @ K_kpe.T)`
# - It computes the attention output and the 2-based log-sum-exp (LSE) of the logits for stable backward passes.
#
# Grid:
# - The grid is 2D: (total_q, num_qo_heads).
# - Each program instance computes the attention output for a single query token and a single head.
#
# Optimization Strategy:
# - Correctness First: The reference `torch.softmax` is a base-e operation. To compute this correctly while using Triton's fast base-2 intrinsics (`tl.exp2`, `tl.log2`), the logits are scaled by `log(2)`. This is based on the identity `softmax_e(x) == softmax_2(x * log(2))`. This ensures numerical alignment with the reference implementation.
# - Two-Pass Stability: A two-pass approach ensures numerical stability for long sequences.
# 1. Pass 1 computes the true base-2 log-sum-exp (LSE) using a stable online algorithm.
# 2. Pass 2 re-computes logits and uses the LSE from Pass 1 to calculate the final attention probabilities and output vector.
# - B200 Optimization: The kernel is tuned with block sizes and parallelization settings (num_warps, num_stages) that are effective on modern architectures. It uses `num_stages` > 1 to pipeline memory loads and compute.
# - Online Softmax: The kernel uses an online (one-pass) softmax algorithm within each pass to handle variable sequence lengths without materializing a large attention matrix.
# - Blocked Computation: All loops over sequence length and head dimensions are blocked to improve data locality.
#
@triton.autotune(
configs=[
triton.Config({'BLOCK_CKV': 64, 'BLOCK_KPE': 64, 'BLOCK_KV': 64}, num_stages=3, num_warps=4),
triton.Config({'BLOCK_CKV': 64, 'BLOCK_KPE': 64, 'BLOCK_KV': 128}, num_stages=2, num_warps=4),
triton.Config({'BLOCK_CKV': 128, 'BLOCK_KPE': 64, 'BLOCK_KV': 64}, num_stages=2, num_warps=4),
triton.Config({'BLOCK_CKV': 64, 'BLOCK_KPE': 64, 'BLOCK_KV': 64}, num_stages=4, num_warps=8),
triton.Config({'BLOCK_CKV': 32, 'BLOCK_KPE': 32, 'BLOCK_KV': 128}, num_stages=3, num_warps=4),
],
key=['total_q', 'num_kv_indices'],
)
@triton.jit
def _mla_paged_prefill_causal_h16_ckv512_kpe64_ps1_kernel(
# Inputs
q_nope_ptr, q_pe_ptr, ckv_cache_ptr, kpe_cache_ptr,
qo_indptr_ptr, kv_indptr_ptr, kv_indices_ptr, q_to_batch_idx_ptr,
sm_scale,
# Outputs
output_ptr, lse_ptr,
# Strides
stride_q_total_q, stride_q_num_heads, stride_q_head_dim_ckv,
stride_qpe_total_q, stride_qpe_num_heads, stride_qpe_head_dim_kpe,
stride_ckv_pages, stride_ckv_page_size, stride_ckv_head_dim,
stride_kpe_pages, stride_kpe_page_size, stride_kpe_head_dim,
stride_out_total_q, stride_out_num_heads, stride_out_head_dim,
stride_lse_total_q, stride_lse_num_heads,
# Axes
total_q: tl.constexpr,
num_pages: tl.constexpr,
len_indptr: tl.constexpr,
num_kv_indices: tl.constexpr,
# Constants
NUM_QO_HEADS: tl.constexpr,
HEAD_DIM_CKV: tl.constexpr,
HEAD_DIM_KPE: tl.constexpr,
PAGE_SIZE: tl.constexpr,
LOG2_E: tl.constexpr,
# Autotune configs
BLOCK_CKV: tl.constexpr,
BLOCK_KPE: tl.constexpr,
BLOCK_KV: tl.constexpr,
):
# =========================================================================
# 1. Program and Grid Setup
# =========================================================================
pid_qt = tl.program_id(0)
pid_h = tl.program_id(1)
if pid_qt >= total_q:
return
# =========================================================================
# 2. Determine Sequence Boundaries
# =========================================================================
batch_idx = tl.load(q_to_batch_idx_ptr + pid_qt)
q_start = tl.load(qo_indptr_ptr + batch_idx)
kv_pages_start = tl.load(kv_indptr_ptr + batch_idx)
kv_pages_end = tl.load(kv_indptr_ptr + batch_idx + 1)
kv_len = kv_pages_end - kv_pages_start
if kv_len == 0:
out_ptr_base = output_ptr + pid_qt * stride_out_total_q + pid_h * stride_out_num_heads
offs_dh = tl.arange(0, BLOCK_CKV)
# Iterate over output head dim to zero out the full vector
for ckv_off in range(0, HEAD_DIM_CKV, BLOCK_CKV):
mask = (ckv_off + offs_dh) < HEAD_DIM_CKV
tl.store(out_ptr_base + ckv_off + offs_dh, tl.zeros((BLOCK_CKV,), dtype=tl.bfloat16), mask=mask)
lse_val_ptr = lse_ptr + pid_qt * stride_lse_total_q + pid_h * stride_lse_num_heads
tl.store(lse_val_ptr, -float('inf'))
return
q_end = tl.load(qo_indptr_ptr + batch_idx + 1)
q_len = q_end - q_start
prefix_len = kv_len - q_len
q_idx_in_seq = pid_qt - q_start
abs_pos_q = prefix_len + q_idx_in_seq
q_nope_offset = pid_qt * stride_q_total_q + pid_h * stride_q_num_heads
q_pe_offset = pid_qt * stride_qpe_total_q + pid_h * stride_qpe_num_heads
# =========================================================================
# 3. Pass 1: Compute LSE (Log-Sum-Exp)
# =========================================================================
m_i = -float("inf")
l_i = 0.0
for kv_block_start in range(0, kv_len, BLOCK_KV):
offs_kv_indices = kv_pages_start + kv_block_start + tl.arange(0, BLOCK_KV)
mask_kv_indices = offs_kv_indices < kv_pages_end
page_indices = tl.load(kv_indices_ptr + offs_kv_indices, mask=mask_kv_indices, other=0)
logits = tl.zeros([BLOCK_KV], dtype=tl.float32)
# CKV component
for ckv_off in range(0, HEAD_DIM_CKV, BLOCK_CKV):
offs_d_ckv = ckv_off + tl.arange(0, BLOCK_CKV)
mask_d_ckv = offs_d_ckv < HEAD_DIM_CKV
q_nope_fragment = tl.load(q_nope_ptr + q_nope_offset + offs_d_ckv, mask=mask_d_ckv, other=0.0).to(tl.float32)
k_ckv = tl.load(ckv_cache_ptr + page_indices[:, None] * stride_ckv_pages + offs_d_ckv[None, :],
mask=mask_kv_indices[:, None] & mask_d_ckv[None, :], other=0.0).to(tl.float32)
logits += tl.sum(q_nope_fragment[None, :] * k_ckv, axis=1)
# KPE component
for kpe_off in range(0, HEAD_DIM_KPE, BLOCK_KPE):
offs_d_kpe = kpe_off + tl.arange(0, BLOCK_KPE)
mask_d_kpe = offs_d_kpe < HEAD_DIM_KPE
q_pe_fragment = tl.load(q_pe_ptr + q_pe_offset + offs_d_kpe, mask=mask_d_kpe, other=0.0).to(tl.float32)
k_kpe = tl.load(kpe_cache_ptr + page_indices[:, None] * stride_kpe_pages + offs_d_kpe[None, :],
mask=mask_kv_indices[:, None] & mask_d_kpe[None, :], other=0.0).to(tl.float32)
logits += tl.sum(q_pe_fragment[None, :] * k_kpe, axis=1)
logits *= sm_scale
# Scale logits by log(2) to compute base-e softmax using base-2 instructions.
# softmax_e(x) == softmax_2(x * log2(e))
logits *= LOG2_E
kv_seq_indices = kv_block_start + tl.arange(0, BLOCK_KV)
causal_mask = kv_seq_indices <= abs_pos_q
final_mask = mask_kv_indices & causal_mask
logits = tl.where(final_mask, logits, -float("inf"))
m_i_new = tl.maximum(m_i, tl.max(logits, axis=0))
p = tl.exp2(logits - m_i_new)
l_i_new = tl.exp2(m_i - m_i_new) * l_i + tl.sum(p, axis=0)
m_i = m_i_new
l_i = l_i_new
# Final LSE is log2(sum(exp(original_logits * sm_scale)))
lse_val = m_i + tl.log2(l_i)
lse_val_ptr = lse_ptr + pid_qt * stride_lse_total_q + pid_h * stride_lse_num_heads
tl.store(lse_val_ptr, lse_val)
# =========================================================================
# 4. Pass 2: Compute Attention Output
# =========================================================================
out_ptr_base = output_ptr + pid_qt * stride_out_total_q + pid_h * stride_out_num_heads
# This pass iterates over the output dimension to keep the accumulator in registers.
for ckv_out_offset in range(0, HEAD_DIM_CKV, BLOCK_CKV):
acc = tl.zeros([BLOCK_CKV], dtype=tl.float32)
# Loop over KV sequence again
for kv_block_start in range(0, kv_len, BLOCK_KV):
offs_kv_indices = kv_pages_start + kv_block_start + tl.arange(0, BLOCK_KV)
mask_kv_indices = offs_kv_indices < kv_pages_end
page_indices = tl.load(kv_indices_ptr + offs_kv_indices, mask=mask_kv_indices, other=0)
# Re-compute logits
logits = tl.zeros([BLOCK_KV], dtype=tl.float32)
for ckv_off in range(0, HEAD_DIM_CKV, BLOCK_CKV):
offs_d_ckv = ckv_off + tl.arange(0, BLOCK_CKV)
mask_d_ckv = offs_d_ckv < HEAD_DIM_CKV
q_nope_fragment = tl.load(q_nope_ptr + q_nope_offset + offs_d_ckv, mask=mask_d_ckv, other=0.0).to(tl.float32)
k_ckv = tl.load(ckv_cache_ptr + page_indices[:, None] * stride_ckv_pages + offs_d_ckv[None, :],
mask=mask_kv_indices[:, None] & mask_d_ckv[None, :], other=0.0).to(tl.float32)
logits += tl.sum(q_nope_fragment[None, :] * k_ckv, axis=1)
for kpe_off in range(0, HEAD_DIM_KPE, BLOCK_KPE):
offs_d_kpe = kpe_off + tl.arange(0, BLOCK_KPE)
mask_d_kpe = offs_d_kpe < HEAD_DIM_KPE
q_pe_fragment = tl.load(q_pe_ptr + q_pe_offset + offs_d_kpe, mask=mask_d_kpe, other=0.0).to(tl.float32)
k_kpe = tl.load(kpe_cache_ptr + page_indices[:, None] * stride_kpe_pages + offs_d_kpe[None, :],
mask=mask_kv_indices[:, None] & mask_d_kpe[None, :], other=0.0).to(tl.float32)
logits += tl.sum(q_pe_fragment[None, :] * k_kpe, axis=1)
logits *= sm_scale
logits *= LOG2_E # Re-apply scaling for base-2 probability calculation
kv_seq_indices = kv_block_start + tl.arange(0, BLOCK_KV)
causal_mask = kv_seq_indices <= abs_pos_q
final_mask = mask_kv_indices & causal_mask
logits = tl.where(final_mask, logits, -float("inf"))
# Compute attention probabilities using the final LSE from Pass 1
p = tl.exp2(logits - lse_val)
# Load V block for the current output slice and update accumulator
offs_v_ckv = ckv_out_offset + tl.arange(0, BLOCK_CKV)
mask_v_ckv = offs_v_ckv < HEAD_DIM_CKV
v_ckv = tl.load(ckv_cache_ptr + page_indices[:, None] * stride_ckv_pages + offs_v_ckv[None, :],
mask=mask_kv_indices[:, None] & mask_v_ckv[None, :], other=0.0)
p = p.to(v_ckv.dtype)
acc += tl.sum(p[:, None] * v_ckv, axis=0)
# Store this block of the output vector
offs_out = ckv_out_offset + tl.arange(0, BLOCK_CKV)
mask_out = offs_out < HEAD_DIM_CKV
tl.store(out_ptr_base + offs_out, acc.to(tl.bfloat16), mask=mask_out)
def _get_sig_bound_args(fn, args, kwargs):
"""Binds `args` and `kwargs` to the signature of `fn`."""
sig = inspect.signature(fn)
bound_args = sig.bind(*args, **kwargs)
bound_args.apply_defaults()
return bound_args.arguments
def _forward(q_nope, q_pe, ckv_cache, kpe_cache, qo_indptr, kv_indptr, kv_indices, sm_scale):
# Shape checks and constants
total_q, num_qo_heads, head_dim_ckv = q_nope.shape
head_dim_kpe = q_pe.shape[-1]
num_pages, page_size, _ = ckv_cache.shape
len_indptr = qo_indptr.shape[0]
num_kv_indices = kv_indices.shape[0]
batch_size = len_indptr - 1
# Assertions for fixed dimensions
assert num_qo_heads == 16, f"Expected num_qo_heads=16, got {num_qo_heads}"
assert head_dim_ckv == 512, f"Expected head_dim_ckv=512, got {head_dim_ckv}"
assert head_dim_kpe == 64, f"Expected head_dim_kpe=64, got {head_dim_kpe}"
assert page_size == 1, f"Expected page_size=1, got {page_size}"
# Create output tensors
output = torch.empty_like(q_nope)
lse = torch.empty((total_q, num_qo_heads), dtype=torch.float32, device=q_nope.device)
# Pre-compute a mapping from query token index to its batch index
q_to_batch_idx = torch.zeros(total_q, dtype=torch.int32, device=q_nope.device)
if total_q > 0 and batch_size > 0:
q_starts = qo_indptr[:-1].long()
q_ends = qo_indptr[1:].long()
for i in range(batch_size):
q_to_batch_idx[q_starts[i]:q_ends[i]] = i
# Grid for kernel launch
grid = (total_q, num_qo_heads)
# Call the Triton kernel
_mla_paged_prefill_causal_h16_ckv512_kpe64_ps1_kernel[grid](
q_nope, q_pe, ckv_cache, kpe_cache,
qo_indptr, kv_indptr, kv_indices, q_to_batch_idx,
sm_scale,
output, lse,
# Strides
q_nope.stride(0), q_nope.stride(1), q_nope.stride(2),
q_pe.stride(0), q_pe.stride(1), q_pe.stride(2),
ckv_cache.stride(0), ckv_cache.stride(1), ckv_cache.stride(2),
kpe_cache.stride(0), kpe_cache.stride(1), kpe_cache.stride(2),
output.stride(0), output.stride(1), output.stride(2),
lse.stride(0), lse.stride(1),
# Axes
total_q, num_pages, len_indptr, num_kv_indices,
# Constants
NUM_QO_HEADS=num_qo_heads,
HEAD_DIM_CKV=head_dim_ckv,
HEAD_DIM_KPE=head_dim_kpe,
PAGE_SIZE=page_size,
LOG2_E=math.log2(math.e),
)
return output, lse
def run(*args, **kwargs):
"""
Wrapper function for the paged prefill attention kernel.
Handles device management and argument binding.
"""
bound_args = _get_sig_bound_args(_forward, args, kwargs)
# Extract tensors from bound arguments
input_tensors_names = ['q_nope', 'q_pe', 'ckv_cache', 'kpe_cache', 'qo_indptr', 'kv_indptr', 'kv_indices']
input_tensors = [bound_args[name] for name in input_tensors_names]
original_device = input_tensors[0].device
is_cpu = original_device.type == 'cpu'
if is_cpu:
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available, but required for Triton kernel execution from CPU tensors.")
gpu_tensors = [t.cuda() for t in input_tensors]
for name, tensor in zip(input_tensors_names, gpu_tensors):
bound_args[name] = tensor
elif not torch.cuda.is_available():
raise RuntimeError("CUDA is not available, but input tensors are on a CUDA device.")
# Execute Forward Pass with potentially moved tensors
output, lse = _forward(**bound_args)
# Restore Original Device if necessary
if is_cpu:
output = output.to(original_device)
lse = lse.to(original_device)
return output, lsescrolls · 318 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON