Skip to content
KernelIndex
Search⌘K

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
NVIDIA B200
162.4µs
#4 of 6
2025-10-21

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 = 8num_warps = 8
online-softmaxm_new = tl.maximum(m, tl.max(s, axis=0))
tile-n = 128BLOCK_N = 128

Kernel 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