Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / tritonpr9imz

gemini-2.5-pro_triton_pr9imz · gemini-2.5-pro · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

No package. Vendor the mirrored source: 264 lines, Apache-2.0, pinned at da91508.

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-pr9imz?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 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=15 · num_kv_indices=14
NVIDIA B200
25.0µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=82 · num_kv_indices=65
NVIDIA B200
25.3µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=67 · num_kv_indices=50
NVIDIA B200
25.4µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=18 · num_kv_indices=2
NVIDIA B200
25.5µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=71 · num_kv_indices=54
NVIDIA B200
26.2µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1220 · num_kv_indices=1193
NVIDIA B200
26.3µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=11 · num_kv_indices=10
NVIDIA B200
26.6µs
#6 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=59 · num_kv_indices=42
NVIDIA B200
26.7µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=9316 · num_kv_indices=73
NVIDIA B200
27.2µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1396 · num_kv_indices=1369
NVIDIA B200
27.4µs
#2 of 7
2025-10-16
Show all 48 measurements ›
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=63 · num_kv_indices=46
NVIDIA B200
27.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=55 · num_kv_indices=38
NVIDIA B200
27.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=9341 · num_kv_indices=98
NVIDIA B200
27.9µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=74 · num_kv_indices=57
NVIDIA B200
28.1µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=78 · num_kv_indices=61
NVIDIA B200
29.0µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=406 · num_kv_indices=356
NVIDIA B200
29.5µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=597 · num_kv_indices=547
NVIDIA B200
29.5µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1556 · num_kv_indices=1529
NVIDIA B200
29.8µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2372 · num_kv_indices=2345
NVIDIA B200
31.6µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1069 · num_kv_indices=1034
NVIDIA B200
32.5µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2708 · num_kv_indices=2681
NVIDIA B200
32.8µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2052 · num_kv_indices=2025
NVIDIA B200
33.2µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2548 · num_kv_indices=2521
NVIDIA B200
34.3µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2868 · num_kv_indices=2841
NVIDIA B200
34.8µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=3044 · num_kv_indices=3017
NVIDIA B200
36.9µs
#2 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1892 · num_kv_indices=1865
NVIDIA B200
42.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=2212 · num_kv_indices=2185
NVIDIA B200
43.2µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [1, 32, 128] · num_pages=223 · num_kv_indices=173
NVIDIA B200
46.1µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=1732 · num_kv_indices=1705
NVIDIA B200
66.4µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=4333 · num_kv_indices=4298
NVIDIA B200
101.4µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=81390 · num_kv_indices=12942
NVIDIA B200
116.2µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [16, 32, 128] · num_pages=30163 · num_kv_indices=20911
NVIDIA B200
156.4µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=60071 · num_kv_indices=50902
NVIDIA B200
252.7µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=62605 · num_kv_indices=53334
NVIDIA B200
258.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=64205 · num_kv_indices=54934
NVIDIA B200
261.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65421 · num_kv_indices=56150
NVIDIA B200
264.1µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65805 · num_kv_indices=56534
NVIDIA B200
265.1µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67021 · num_kv_indices=57750
NVIDIA B200
269.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67405 · num_kv_indices=58134
NVIDIA B200
270.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=67853 · num_kv_indices=58582
NVIDIA B200
274.7µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63053 · num_kv_indices=53782
NVIDIA B200
285.3µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=65037 · num_kv_indices=55766
NVIDIA B200
293.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=66253 · num_kv_indices=56982
NVIDIA B200
299.3µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=68237 · num_kv_indices=58966
NVIDIA B200
328.6µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63437 · num_kv_indices=54166
NVIDIA B200
461.9µs
#4 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=66637 · num_kv_indices=57366
NVIDIA B200
480.9µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=63821 · num_kv_indices=54550
NVIDIA B200
893.5µs
#5 of 7
2025-10-16
GQA paged decode h32 kv8 d128 ps1bf16 · [64, 32, 128] · num_pages=64653 · num_kv_indices=55382
NVIDIA B200
903.7µs
#5 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:e5d6283af3f6143bf93e3bc93f551732f18d50a9d53538cf7fead3802f18b330
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 = 4triton.Config({'BLOCK_KV_LEN': 16}, num_warps=4),

Kernel source

main.py264 lines
import torch
import triton
import triton.language as tl
import math


@triton.autotune(
    configs=[
        triton.Config({'BLOCK_KV_LEN': 16}, num_warps=4),
        triton.Config({'BLOCK_KV_LEN': 32}, num_warps=4),
        triton.Config({'BLOCK_KV_LEN': 64}, num_warps=4),
        triton.Config({'BLOCK_KV_LEN': 128}, num_warps=4),
        triton.Config({'BLOCK_KV_LEN': 256}, num_warps=4),
        triton.Config({'BLOCK_KV_LEN': 16}, num_warps=8),
        triton.Config({'BLOCK_KV_LEN': 32}, num_warps=8),
        triton.Config({'BLOCK_KV_LEN': 64}, num_warps=8),
        triton.Config({'BLOCK_KV_LEN': 128}, num_warps=8),
        triton.Config({'BLOCK_KV_LEN': 256}, num_warps=8),
    ],
    key=['HEAD_DIM'],
)
@triton.jit
def gqa_paged_decode_h32_kv8_d128_ps1_kernel(
    # Pointers to Tensors
    Q_ptr, K_cache_ptr, V_cache_ptr,
    kv_indptr_ptr, kv_indices_ptr,
    sm_scale,
    Output_ptr, LSE_ptr,
    # Stride Info
    stride_q_bs, stride_q_h,
    stride_k_num_pages, stride_k_ps, stride_k_h,
    stride_v_num_pages, stride_v_ps, stride_v_h,
    stride_o_bs, stride_o_h,
    stride_lse_bs,
    # Compile-time Constants
    NUM_QO_HEADS: tl.constexpr,
    NUM_KV_HEADS: tl.constexpr,
    HEAD_DIM: tl.constexpr,
    GQA_RATIO: tl.constexpr,
    PAGE_SIZE: tl.constexpr,
    BLOCK_D: tl.constexpr,
    BLOCK_KV_LEN: tl.constexpr,
):
    """
    Triton kernel for GQA paged decode attention.

    This kernel computes attention for a batch of query vectors against their
    corresponding key/value history, which is stored in a paged cache. Each
    program instance handles one query head for one sequence in the batch.

    Grid: (batch_size, num_qo_heads)
    - pid_b (program_id 0): batch index
    - pid_h (program_id 1): query head index

    Key Optimizations for B200:
    - Processes the variable-length KV sequence in fixed-size blocks (BLOCK_KV_LEN)
      to increase arithmetic intensity and hide memory latency from gathers.
    - Online softmax algorithm is used to compute attention scores and output
      in a single pass over the KV sequence, avoiding materialization of the
      full attention matrix.
    - For this decode-style kernel (one query vector per program), the Q@K.T and P@V
      operations are GEMV-like. Since tl.dot has minimum shape requirements (e.g., M>=16)
      for Tensor Core usage that are not met by a single vector, these operations
      are implemented using efficient element-wise operations and reductions.
    - Autotuning is enabled for BLOCK_KV_LEN and num_warps to find the optimal
      configuration for the target hardware.
    """
    # 1. Get program IDs for batch and query head
    pid_b = tl.program_id(0)
    pid_h = tl.program_id(1)

    # 2. Determine KV sequence length and handle empty sequences
    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

    # Early exit for sequences with no KV history
    if seq_len == 0:
        offs_d = tl.arange(0, BLOCK_D)
        out_ptr = Output_ptr + pid_b * stride_o_bs + pid_h * stride_o_h + offs_d
        lse_ptr = LSE_ptr + pid_b * stride_lse_bs + pid_h

        tl.store(out_ptr, tl.zeros([BLOCK_D], dtype=tl.bfloat16), mask=offs_d < BLOCK_D)
        tl.store(lse_ptr, -float('inf'))
        return

    # 3. Load query vector
    offs_d = tl.arange(0, BLOCK_D)
    q_ptr = Q_ptr + pid_b * stride_q_bs + pid_h * stride_q_h + offs_d
    q = tl.load(q_ptr, mask=offs_d < BLOCK_D).to(tl.float32)[None, :]  # Shape: [1, BLOCK_D]

    # 4. Initialize accumulators for online softmax
    # FIX: Use scalar accumulators for m_i and l_i to avoid shape errors
    # during the final tl.store operation for the LSE scalar.
    acc_o = tl.zeros([BLOCK_D], dtype=tl.float32)
    m_i = -float('inf')
    l_i = 0.0

    # 5. Determine the corresponding KV head for GQA
    kv_head_idx = pid_h // GQA_RATIO

    # 6. Loop over the KV sequence in blocks
    for offset in range(0, seq_len, BLOCK_KV_LEN):
        # a. Create masks and pointers for the current block
        offs_kv_block = offset + tl.arange(0, BLOCK_KV_LEN)
        mask_kv_block = offs_kv_block < seq_len
        indices_ptr = kv_indices_ptr + page_start + offs_kv_block

        # b. Gather page indices for K and V caches
        page_indices = tl.load(indices_ptr, mask=mask_kv_block, other=0)

        # c. Gather K vectors for the block
        offs_k_h = kv_head_idx * stride_k_h
        k_ptrs = K_cache_ptr + page_indices[:, None] * stride_k_num_pages + offs_k_h + offs_d[None, :]
        k = tl.load(k_ptrs, mask=mask_kv_block[:, None] & (offs_d[None, :] < BLOCK_D), other=0.0)

        # d. Compute scores S = Q @ K.T
        s_block = tl.sum(q * k.to(tl.float32), axis=1) * sm_scale
        # FIX: Keep scores as a 1D tensor [BLOCK_KV_LEN] for scalar reduction.
        s = tl.where(mask_kv_block, s_block, -float('inf'))

        # e. Online softmax update (with scalar state)
        m_block_max = tl.max(s, axis=0)
        m_curr = tl.maximum(m_i, m_block_max)
        p = tl.exp(s - m_curr)
        l_i_exp = tl.exp(m_i - m_curr)
        l_curr = l_i_exp * l_i + tl.sum(p, axis=0)

        # f. Gather V vectors for the block
        offs_v_h = kv_head_idx * stride_v_h
        v_ptrs = V_cache_ptr + page_indices[:, None] * stride_v_num_pages + offs_v_h + offs_d[None, :]
        v = tl.load(v_ptrs, mask=mask_kv_block[:, None] & (offs_d[None, :] < BLOCK_D), other=0.0)

        # g. Update output accumulator
        acc_o = acc_o * l_i_exp
        # FIX: Reshape 1D p to [BLOCK_KV_LEN, 1] for broadcasted matmul-like update.
        update_o = tl.sum(p[:, None] * v.to(tl.float32), axis=0)
        acc_o += update_o

        # h. Update state for the next iteration
        m_i = m_curr
        l_i = l_curr

    # 7. Finalize and store results
    l_i_safe = tl.where(l_i == 0, 1.0, l_i)
    acc_o = acc_o / l_i_safe

    # Calculate 2-based log-sum-exp
    lse = m_i + tl.log(l_i_safe)
    lse = lse / 0.6931471805599453

    # Store output vector and LSE value
    out_ptr = Output_ptr + pid_b * stride_o_bs + pid_h * stride_o_h + offs_d
    tl.store(out_ptr, acc_o.to(tl.bfloat16), mask=offs_d < BLOCK_D)
    # FIX: Storing the scalar `lse` value to a scalar pointer is now valid.
    tl.store(LSE_ptr + pid_b * stride_lse_bs + pid_h, lse)


def _gqa_paged_decode_h32_kv8_d128_ps1_launcher(q, k_cache, v_cache, kv_indptr, kv_indices, sm_scale):
    """
    Host-side wrapper for the Triton kernel.

    This function handles device management, output tensor allocation, grid
    computation, and kernel invocation. It ensures all tensors are on the
    same CUDA device before launching the kernel and moves the results back
    to the original device of the `q` tensor.
    """
    # 1. Validate inputs and extract dimensions
    assert q.dtype == torch.bfloat16
    assert k_cache.dtype == torch.bfloat16
    assert v_cache.dtype == torch.bfloat16
    assert kv_indptr.dtype == torch.int32
    assert kv_indices.dtype == torch.int32

    batch_size, num_qo_heads, head_dim = q.shape
    num_pages, page_size, num_kv_heads, _ = k_cache.shape

    # Check problem-specific constants
    assert num_qo_heads == 32, f"Expected num_qo_heads=32, got {num_qo_heads}"
    assert num_kv_heads == 8, f"Expected num_kv_heads=8, got {num_kv_heads}"
    assert head_dim == 128, f"Expected head_dim=128, got {head_dim}"
    assert page_size == 1, f"Expected page_size=1, got {page_size}"

    # 2. Set default sm_scale if not provided
    if sm_scale is None:
        sm_scale = 1.0 / math.sqrt(head_dim)

    # 3. Complete device management
    original_q_device = q.device

    # Find a common CUDA device for computation
    target_device = None
    for t in [q, k_cache, v_cache, kv_indptr, kv_indices]:
        if t.is_cuda:
            target_device = t.device
            break
    if target_device is None:
        if torch.cuda.is_available():
            target_device = torch.device("cuda")
        else:
            raise RuntimeError("This Triton kernel requires a CUDA-enabled GPU, but none was found.")

    # Move all tensors to the target CUDA device
    q = q.to(target_device)
    k_cache = k_cache.to(target_device)
    v_cache = v_cache.to(target_device)
    kv_indptr = kv_indptr.to(target_device)
    kv_indices = kv_indices.to(target_device)

    # 4. Allocate output tensors on the target device
    output = torch.empty_like(q)
    lse = torch.empty((batch_size, num_qo_heads), dtype=torch.float32, device=target_device)

    # 5. Set up grid and launch kernel
    grid = (batch_size, num_qo_heads)

    gqa_paged_decode_h32_kv8_d128_ps1_kernel[grid](
        q, k_cache, v_cache,
        kv_indptr, kv_indices,
        float(sm_scale),
        output, lse,
        # Strides
        q.stride(0), q.stride(1),
        k_cache.stride(0), k_cache.stride(1), k_cache.stride(2),
        v_cache.stride(0), v_cache.stride(1), v_cache.stride(2),
        output.stride(0), output.stride(1),
        lse.stride(0),
        # Constants
        NUM_QO_HEADS=num_qo_heads,
        NUM_KV_HEADS=num_kv_heads,
        HEAD_DIM=head_dim,
        GQA_RATIO=num_qo_heads // num_kv_heads,
        PAGE_SIZE=page_size,
        BLOCK_D=head_dim,
    )

    # 6. Move results back to the original device of `q`
    output = output.to(original_q_device)
    lse = lse.to(original_q_device)

    return output, lse

def run(*args, **kwargs):
    """
    Public entry point for the GQA paged decode attention kernel.

    This function acts as a flexible interface, accepting both positional and
    keyword arguments and forwarding them to the core launcher function.

    Args:
        q (torch.Tensor): Query tensor of shape [batch_size, 32, 128] and dtype bfloat16.
        k_cache (torch.Tensor): Key cache of shape [num_pages, 1, 8, 128] and dtype bfloat16.
        v_cache (torch.Tensor): Value cache of shape [num_pages, 1, 8, 128] and dtype bfloat16.
        kv_indptr (torch.Tensor): KV page offsets of shape [batch_size + 1] and dtype int32.
        kv_indices (torch.Tensor): Page IDs of shape [num_kv_indices] and dtype int32.
        sm_scale (float, optional): Softmax scale. Defaults to 1/sqrt(head_dim).

    Returns:
        Tuple[torch.Tensor, torch.Tensor]:
            - output: The attention output tensor of shape [batch_size, 32, 128] and dtype bfloat16.
            - lse: The log-sum-exp of attention logits (base 2) of shape [batch_size, 32] and dtype float32.
    """
    return _gqa_paged_decode_h32_kv8_d128_ps1_launcher(*args, **kwargs)
scrolls · 264 lines total

Source code from FlashInfer-Bench (flashinfer-ai/flashinfer-trace) · Apache-2.0

Best evidence level for this revision: reported

JSON