Skip to content
KernelIndex
Search⌘K

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 = 4triton.Config({'BLOCK_CKV': 64, 'BLOCK_KPE': 64, 'BLOCK_KV': 64}, num_stages=3, num_warps=4),
online-softmaxm_i_new = tl.maximum(m_i, tl.max(logits, axis=0))
stages = 3triton.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, lse
scrolls · 318 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON