Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / tritonh7ykt0

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-h7ykt0?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 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9332 · num_kv_indices=87
NVIDIA B200
64.9µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=10 · num_kv_indices=9
NVIDIA B200
66.2µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9317 · num_kv_indices=72
NVIDIA B200
66.6µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=15 · num_kv_indices=14
NVIDIA B200
66.8µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=57 · num_kv_indices=40
NVIDIA B200
66.8µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=67 · num_kv_indices=50
NVIDIA B200
66.9µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=12 · num_kv_indices=11
NVIDIA B200
67.3µs
#6 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=81 · num_kv_indices=64
NVIDIA B200
67.4µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=8 · num_kv_indices=7
NVIDIA B200
67.5µs
#5 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=17 · num_kv_indices=2
NVIDIA B200
67.5µs
#5 of 7
2025-10-16
Show all 48 measurements ›
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=9347 · num_kv_indices=102
NVIDIA B200
67.5µs
#2 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=191 · num_kv_indices=141
NVIDIA B200
78.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=302 · num_kv_indices=252
NVIDIA B200
79.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=1070 · num_kv_indices=1020
NVIDIA B200
93.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=412 · num_kv_indices=362
NVIDIA B200
94.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=2158 · num_kv_indices=2108
NVIDIA B200
103.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=3246 · num_kv_indices=3196
NVIDIA B200
104.3µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=486 · num_kv_indices=436
NVIDIA B200
106.7µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [1, 32, 128] · num_pages=596 · num_kv_indices=546
NVIDIA B200
117.1µs
#4 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=5422 · num_kv_indices=5372
NVIDIA B200
117.4µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=4334 · num_kv_indices=4284
NVIDIA B200
119.1µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=6510 · num_kv_indices=6460
NVIDIA B200
127.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=7598 · num_kv_indices=7548
NVIDIA B200
129.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=8686 · num_kv_indices=8636
NVIDIA B200
144.8µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=24732 · num_kv_indices=15463
NVIDIA B200
373.3µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=28831 · num_kv_indices=28815
NVIDIA B200
374.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=31007 · num_kv_indices=30991
NVIDIA B200
374.4µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=26908 · num_kv_indices=17639
NVIDIA B200
386.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=33183 · num_kv_indices=33167
NVIDIA B200
386.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=25820 · num_kv_indices=16551
NVIDIA B200
386.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=35359 · num_kv_indices=35343
NVIDIA B200
387.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=39711 · num_kv_indices=39695
NVIDIA B200
389.7µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=27996 · num_kv_indices=18727
NVIDIA B200
398.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=41887 · num_kv_indices=41871
NVIDIA B200
400.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=44063 · num_kv_indices=44047
NVIDIA B200
402.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=46303 · num_kv_indices=46287
NVIDIA B200
402.7µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=48479 · num_kv_indices=48463
NVIDIA B200
404.5µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=29084 · num_kv_indices=19815
NVIDIA B200
404.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=30172 · num_kv_indices=20903
NVIDIA B200
411.8µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=52831 · num_kv_indices=52815
NVIDIA B200
414.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=50655 · num_kv_indices=50639
NVIDIA B200
416.1µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=55007 · num_kv_indices=54991
NVIDIA B200
417.2µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=31260 · num_kv_indices=21991
NVIDIA B200
425.6µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=61535 · num_kv_indices=61519
NVIDIA B200
426.7µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [16, 32, 128] · num_pages=32348 · num_kv_indices=23079
NVIDIA B200
427.9µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=59359 · num_kv_indices=59343
NVIDIA B200
428.3µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=57183 · num_kv_indices=57167
NVIDIA B200
430.0µs
#3 of 7
2025-10-16
GQA paged decode h32 kv4 d128 ps1bf16 · [64, 32, 128] · num_pages=37535 · num_kv_indices=37519
NVIDIA B200
430.1µs
#3 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:6ff3d45887bd6fffc9365b005926d1286c2456c122012f6932ad59d05a47a9eb
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.

mmas_ij_subset = tl.dot(q_subset, tl.trans(k.to(tl.float32)))
num-warps = 4num_warps = 4
online-softmaxm_new = tl.maximum(m_i, m_ij)
stages = 3num_stages = 3
tile-n = 128BLOCK_N = 128

Kernel source

main.py268 lines
import torch
import triton
import triton.language as tl
import math
from typing import Optional

# Constant for converting natural log to base-2 log.
INV_LOG_2 = 1.0 / math.log(2.0)


@triton.jit
def gqa_paged_decode_h32_kv4_d128_ps1_kernel(
    # Pointers to Tensors
    q_ptr, k_cache_ptr, v_cache_ptr,
    kv_indptr_ptr, kv_indices_ptr,
    sm_scale,
    lse_ptr, output_ptr,
    # Stride information
    stride_q_bs, stride_q_h,
    stride_k_page, stride_k_head,
    stride_v_page, stride_v_head,
    stride_out_bs, stride_out_h,
    stride_lse_bs, stride_lse_h,
    # GQA parameters
    gqa_ratio: tl.constexpr,
    # Meta-parameters
    NUM_QO_HEADS: tl.constexpr,
    NUM_KV_HEADS: tl.constexpr,
    HEAD_DIM: tl.constexpr,
    BLOCK_N: tl.constexpr,
    INV_LOG_2: tl.constexpr,
):
    """
    Triton kernel for paged GQA decode.
    Each program computes the attention output for ALL query heads of one sequence.
    This is done to satisfy the M >= 16 constraint of tl.dot on B200/H100.
    """
    # 1. Get program ID for the batch dimension
    pid_b = tl.program_id(0)

    # 2. Load sequence information from kv_indptr
    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

    # 3. Define offsets for head and dimension axes
    offs_h = tl.arange(0, NUM_QO_HEADS)
    offs_d = tl.arange(0, HEAD_DIM)

    # 4. Early exit for sequences with no KV cache
    if seq_len == 0:
        # Store zero output
        output_offset = pid_b * stride_out_bs + offs_h[:, None] * stride_out_h + offs_d[None, :]
        tl.store(output_ptr + output_offset, tl.zeros([NUM_QO_HEADS, HEAD_DIM], dtype=tl.bfloat16))
        
        # Store -inf LSE
        lse_offset = pid_b * stride_lse_bs + offs_h * stride_lse_h
        tl.store(lse_ptr + lse_offset, tl.full([NUM_QO_HEADS], -float('inf'), dtype=tl.float32))
        return

    # 5. Load Q matrix for all heads of the current sequence
    q_offset = pid_b * stride_q_bs + offs_h[:, None] * stride_q_h + offs_d[None, :]
    q = tl.load(q_ptr + q_offset).to(tl.float32)

    # 6. Initialize accumulators for online softmax (one per head)
    # Shapes are [NUM_QO_HEADS, 1] for broadcasting with scores [NUM_QO_HEADS, BLOCK_N]
    acc_o = tl.zeros([NUM_QO_HEADS, HEAD_DIM], dtype=tl.float32)
    m_i = tl.full([NUM_QO_HEADS, 1], -float('inf'), dtype=tl.float32)
    l_i = tl.zeros([NUM_QO_HEADS, 1], dtype=tl.float32)

    # 7. Determine the corresponding KV head index for each Q head
    kv_head_indices = offs_h // gqa_ratio

    # 8. Main loop over the KV sequence length in blocks of BLOCK_N
    offs_n = tl.arange(0, BLOCK_N)
    kv_indices_base_ptr = kv_indices_ptr + page_start
    
    num_blocks = tl.cdiv(seq_len, BLOCK_N)
    for block_idx in range(num_blocks):
        # a. Compute offsets and masks for the current block
        current_block_start = block_idx * BLOCK_N
        kv_indices_offs = current_block_start + offs_n
        kv_mask = kv_indices_offs < seq_len
        page_ids = tl.load(kv_indices_base_ptr + kv_indices_offs, mask=kv_mask, other=0)

        # b. Compute scores S = Q @ K.T
        # We iterate over each KV head, compute scores for the corresponding Q heads,
        # and accumulate the results.
        s_ij = tl.zeros([NUM_QO_HEADS, BLOCK_N], dtype=tl.float32)
        for kv_h_idx in range(NUM_KV_HEADS):
            # Mask to select Q heads corresponding to the current KV head
            q_mask = (kv_head_indices == kv_h_idx)
            q_subset = tl.where(q_mask[:, None], q, 0.0)
            
            # Gather load K block for the current KV head
            k_ptr = k_cache_ptr + page_ids[:, None] * stride_k_page + \
                    kv_h_idx * stride_k_head + offs_d[None, :]
            k = tl.load(k_ptr, mask=kv_mask[:, None], other=0.0)

            # Compute scores for this subset of heads
            s_ij_subset = tl.dot(q_subset, tl.trans(k.to(tl.float32)))
            s_ij += s_ij_subset

        s_ij *= sm_scale
        s_ij = tl.where(kv_mask[None, :], s_ij, -float('inf'))

        # c. Update online softmax statistics (m_i, l_i)
        m_ij = tl.max(s_ij, 1)[:, None]
        m_new = tl.maximum(m_i, m_ij)
        
        alpha = tl.exp(m_i - m_new)
        beta = tl.exp(s_ij - m_new)
        
        l_i_update = tl.sum(beta, 1)[:, None]
        l_i = l_i * alpha + l_i_update
        
        # d. Compute P = softmax(S) and update output accumulator (acc_o)
        p_ij = beta.to(tl.bfloat16)
        
        # Rescale old accumulator
        acc_o = acc_o * alpha

        # Iterate over KV heads again to compute P @ V
        for kv_h_idx in range(NUM_KV_HEADS):
            # Mask to select probabilities for the current KV head
            q_mask = (kv_head_indices == kv_h_idx)
            p_ij_subset = tl.where(q_mask[:, None], p_ij, 0.0)

            # Gather load V block for the current KV head
            v_ptr = v_cache_ptr + page_ids[:, None] * stride_v_page + \
                    kv_h_idx * stride_v_head + offs_d[None, :]
            v = tl.load(v_ptr, mask=kv_mask[:, None], other=0.0)

            # Accumulate P @ V for this subset of heads
            acc_o += tl.dot(p_ij_subset, v, out_dtype=tl.float32)
        
        m_i = m_new

    # 9. Finalize and store output and LSE
    # Rescale accumulator to get the final output vector
    l_i_safe = tl.where(l_i == 0.0, 1.0, l_i)
    o = acc_o / l_i_safe
    
    output_offset = pid_b * stride_out_bs + offs_h[:, None] * stride_out_h + offs_d[None, :]
    tl.store(output_ptr + output_offset, o.to(tl.bfloat16))

    # Compute and store 2-based log-sum-exp
    # The indexing `[:, 0]` is not supported by the Triton compiler in this context.
    # Use tl.ravel to flatten the [NUM_QO_HEADS, 1] tensor to [NUM_QO_HEADS]
    # before storing, which matches the 1D shape of the destination offsets.
    log_lse = m_i + tl.log(l_i)
    lse = tl.ravel(log_lse * INV_LOG_2)
    lse_offset = pid_b * stride_lse_bs + offs_h * stride_lse_h
    tl.store(lse_ptr + lse_offset, lse)


def run(
    q: torch.Tensor,
    k_cache: torch.Tensor,
    v_cache: torch.Tensor,
    kv_indptr: torch.Tensor,
    kv_indices: torch.Tensor,
    sm_scale: Optional[float] = None,
) -> (torch.Tensor, torch.Tensor):
    """
    Wrapper function for the GQA Paged Decode kernel.

    Handles device management, tensor validation, kernel launching, and
    returning results to the original device.

    Args:
        q: Query tensor of shape [batch_size, num_qo_heads, head_dim].
        k_cache: Key cache tensor of shape [num_pages, page_size, num_kv_heads, head_dim].
        v_cache: Value cache tensor of shape [num_pages, page_size, num_kv_heads, head_dim].
        kv_indptr: KV page offsets for each sequence, shape [batch_size + 1].
        kv_indices: Page IDs for KV cache lookups, shape [num_kv_indices].
        sm_scale: Softmax scale factor. Defaults to 1/sqrt(head_dim).

    Returns:
        A tuple containing:
        - output: The attention output tensor of shape [batch_size, num_qo_heads, head_dim].
        - lse: The log-sum-exp of attention logits (base 2), shape [batch_size, num_qo_heads].
    """
    # 1. --- Device Management & Validation ---
    if not torch.cuda.is_available():
        raise RuntimeError("Triton kernel requires a CUDA-enabled device.")

    original_device = q.device
    is_cpu_run = original_device.type == 'cpu'
    
    if is_cpu_run:
        # Move all tensor inputs to the default CUDA device
        device = "cuda"
        q = q.to(device)
        k_cache = k_cache.to(device)
        v_cache = v_cache.to(device)
        kv_indptr = kv_indptr.to(device)
        kv_indices = kv_indices.to(device)
    else:
        # Ensure all tensors are on the same CUDA device
        device = q.device
        for t_name, t in [("k_cache", k_cache), ("v_cache", v_cache), ("kv_indptr", kv_indptr), ("kv_indices", kv_indices)]:
            if t.device != device:
                raise ValueError(f"All input tensors must be on the same device. "
                                 f"Expected {device}, but found '{t_name}' on {t.device}.")

    # 2. --- Shape and Parameter Validation ---
    batch_size, num_qo_heads, head_dim = q.shape
    num_pages, page_size, num_kv_heads, _ = k_cache.shape

    # Constants from spec
    if num_qo_heads != 32: raise ValueError(f"Expected num_qo_heads=32, got {num_qo_heads}")
    if num_kv_heads != 4: raise ValueError(f"Expected num_kv_heads=4, got {num_kv_heads}")
    if head_dim != 128: raise ValueError(f"Expected head_dim=128, got {head_dim}")
    if page_size != 1: raise ValueError(f"Expected page_size=1, got {page_size}")

    # Constraints from spec
    if kv_indptr.shape != (batch_size + 1,):
        raise ValueError(f"Expected kv_indptr shape {(batch_size + 1,)}, got {kv_indptr.shape}")
    if kv_indices.shape[0] != kv_indptr[-1].item():
         raise ValueError(f"Mismatch in total number of KV indices.")

    if sm_scale is None:
        sm_scale = 1.0 / math.sqrt(head_dim)

    # 3. --- Kernel Launch Setup ---
    output = torch.empty_like(q)
    lse = torch.empty((batch_size, num_qo_heads), dtype=torch.float32, device=device)

    gqa_ratio = num_qo_heads // num_kv_heads
    # Each program handles all heads for one batch item
    grid = (batch_size,)

    # Kernel meta-parameters optimized for B200
    BLOCK_N = 128
    num_warps = 4
    num_stages = 3

    # 4. --- Launch Kernel ---
    gqa_paged_decode_h32_kv4_d128_ps1_kernel[grid](
        q, k_cache, v_cache, kv_indptr, kv_indices,
        float(sm_scale),
        lse, output,
        # Strides
        q.stride(0), q.stride(1),
        k_cache.stride(0), k_cache.stride(2),
        v_cache.stride(0), v_cache.stride(2),
        output.stride(0), output.stride(1),
        lse.stride(0), lse.stride(1),
        # GQA parameters
        gqa_ratio=gqa_ratio,
        # Meta-parameters
        NUM_QO_HEADS=num_qo_heads,
        NUM_KV_HEADS=num_kv_heads,
        HEAD_DIM=head_dim,
        BLOCK_N=BLOCK_N,
        INV_LOG_2=INV_LOG_2,
        num_warps=num_warps,
        num_stages=num_stages,
    )

    # 5. --- Return Results ---
    if is_cpu_run:
        # Move results back to the original CPU device
        output = output.to(original_device)
        lse = lse.to(original_device)

    return output, lse
scrolls · 268 lines total

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

Best evidence level for this revision: reported

JSON