Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro_triton_zezbpc

gemini-2.5-pro · triton · Apache-2.0

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-zezbpc?include=source"
interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32

Benchmark evidence

15 measurements across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
GQA ragged prefill causal h32 kv4 d128bf16 · [1, 4, 128] · #7c206f
NVIDIA B200
97.6µs
#4 of 5
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [6, 4, 128] · #6d6644
NVIDIA B200
97.8µs
#4 of 5
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [34, 4, 128] · #816a2c
NVIDIA B200
100.4µs
#2 of 5
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [34, 4, 128] · #641e77
NVIDIA B200
103.9µs
#5 of 20
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [6, 4, 128] · #55a16d
NVIDIA B200
104.1µs
#7 of 10
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [6, 4, 128] · #55a16d
NVIDIA B200
104.6µs
#8 of 10
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [34, 4, 128] · #641e77
NVIDIA B200
104.7µs
#6 of 20
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [1, 4, 128] · #ce8167
NVIDIA B200
105.2µs
#7 of 10
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [1, 4, 128] · #ce8167
NVIDIA B200
105.4µs
#8 of 10
2025-10-20
GQA ragged prefill causal h32 kv4 d128bf16 · [34, 4, 128] · #641e77
NVIDIA B200
106.1µs
#7 of 20
2025-10-20
Show all 15 measurements ›
GQA ragged prefill causal h32 kv4 d128bf16 · [34, 4, 128] · #641e77
NVIDIA B200
106.3µs
#8 of 20
2025-10-20
NVIDIA B200
140.2µs
#2 of 5
2025-10-20
NVIDIA B200
563.0µs
#2 of 5
2025-10-20
NVIDIA B200
35.5ms
#3 of 5
2025-10-20
NVIDIA B200
51.7ms
#3 of 5
2025-10-20

Reported · How evidence levels are derived →

Source and license

sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:6c278eec14ff6845c07c4715a9cbb85d5bcaf49d7b90ecfbd3baf5b1787a1641
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.

tile-n = 64BLOCK_N = 64

Kernel source

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

@triton.jit
def gqa_ragged_prefill_causal_h32_kv4_d128_kernel(
    # Pointers to tensors
    q_ptr, k_ptr, v_ptr,
    qo_indptr_ptr, kv_indptr_ptr, q_to_b_idx_ptr,
    output_ptr, lse_ptr,
    # Scalar
    sm_scale,
    # Strides
    q_stride_tq, q_stride_h,
    k_stride_tk, k_stride_h,
    v_stride_tk, v_stride_h,
    # Other metadata
    total_q,
    # Constants for clarity and performance
    GQA_RATIO: tl.constexpr,
    NUM_QO_HEADS: tl.constexpr,
    # Compile-time constants
    HEAD_DIM: tl.constexpr,
    BLOCK_N: tl.constexpr,
):
    """
    Triton kernel for Grouped-Query Attention on ragged tensors for prefill.
    This kernel is specialized for causal attention with specific head dimensions.
    Each program instance computes the attention output for one query token and one query head.
    """
    # Get program IDs to identify the current query token and head
    global_q_idx = tl.program_id(0)
    h_qo_idx = tl.program_id(1)

    # Find the sequence (batch element) index for the current query token
    b_idx = tl.load(q_to_b_idx_ptr + global_q_idx)
    
    # Load sequence boundaries from indptr tensors
    q_start = tl.load(qo_indptr_ptr + b_idx)
    q_end = tl.load(qo_indptr_ptr + b_idx + 1)
    kv_start = tl.load(kv_indptr_ptr + b_idx)
    kv_end = tl.load(kv_indptr_ptr + b_idx + 1)

    # Calculate causal attention length limit
    q_idx_in_seq = global_q_idx - q_start
    delta = (kv_end - kv_start) - (q_end - q_start)
    max_kv_len = q_idx_in_seq + 1 + delta

    # Initialize accumulators for online softmax
    m_i = -float('inf')
    l_i = 0.0
    acc = tl.zeros([HEAD_DIM], dtype=tl.float32)

    # Determine the corresponding KV head for the current QO head
    h_kv_idx = h_qo_idx // GQA_RATIO

    # Load the query vector
    d_offsets = tl.arange(0, HEAD_DIM)
    q_offset = global_q_idx * q_stride_tq + h_qo_idx * q_stride_h
    q_ptrs = q_ptr + q_offset + d_offsets
    q_vec = tl.load(q_ptrs).to(tl.float32)

    # Loop over the key/value sequence in blocks
    num_n_blocks = (max_kv_len + BLOCK_N - 1) // BLOCK_N
    for block_n_idx in range(num_n_blocks):
        # --- Compute offsets and mask for the current block of K/V ---
        kv_idx_in_seq_start = block_n_idx * BLOCK_N
        n_offsets = kv_idx_in_seq_start + tl.arange(0, BLOCK_N)
        kv_mask = n_offsets < max_kv_len
        global_kv_indices = kv_start + n_offsets

        # --- Load K block ---
        k_offset = global_kv_indices * k_stride_tk + h_kv_idx * k_stride_h
        k_ptrs = k_ptr + k_offset[:, None] + d_offsets[None, :]
        k_block = tl.load(k_ptrs, mask=kv_mask[:, None], other=0.0).to(tl.float32)
        
        # --- Compute S = Q @ K.T ---
        s_block = tl.sum(q_vec[None, :] * k_block, axis=1)
        s_block = s_block * sm_scale
        s_block = tl.where(kv_mask, s_block, -float('inf'))

        # --- Online softmax update ---
        m_i_prev = m_i
        m_i = tl.maximum(m_i, tl.max(s_block, axis=0))
        p = tl.exp(s_block - m_i)
        l_i = l_i * tl.exp(m_i_prev - m_i) + tl.sum(p, axis=0)

        # --- Load V block and update accumulator ---
        v_offset = global_kv_indices * v_stride_tk + h_kv_idx * v_stride_h
        v_ptrs = v_ptr + v_offset[:, None] + d_offsets[None, :]
        v_block = tl.load(v_ptrs, mask=kv_mask[:, None], other=0.0).to(tl.float32)

        # Rescale accumulator before adding new values
        acc = acc * tl.exp(m_i_prev - m_i)

        # FIX: The original tl.dot(p, v_block) caused a compilation error because `p` is 1D
        # while tl.dot requires 2D inputs for matrix multiplication.
        # The correct operation is a weighted sum of value vectors: sum(p[i] * v_block[i]).
        # This is implemented by reshaping p to [BLOCK_N, 1] for broadcasting,
        # multiplying with v_block, and then summing over the block dimension (axis=0).
        acc += tl.sum(p[:, None] * v_block, axis=0)

    # Finalize and store output vector
    # Guard against division by zero if l_i is 0 (e.g., empty sequence)
    o = tl.where(l_i > 0, acc / l_i, 0.0)
    output_offset = global_q_idx * q_stride_tq + h_qo_idx * q_stride_h
    output_ptrs = output_ptr + output_offset + d_offsets
    tl.store(output_ptrs, o.to(tl.bfloat16))

    # Finalize and store log-sum-exp (LSE)
    LOG2_E = 1.4426950408889634  # 1.0 / math.log(2.0)
    # Guard against log(0)
    lse = m_i + tl.log(l_i + 1e-9)
    lse = lse * LOG2_E
    lse_offset = global_q_idx * NUM_QO_HEADS + h_qo_idx
    tl.store(lse_ptr + lse_offset, lse)


def _get_device(*tensors):
    """
    Gets the common device of a list of tensors, handling CPU/CUDA logic.
    """
    devices = {t.device.type for t in tensors if hasattr(t, 'device')}
    if not devices:
        return torch.device('cpu')
    
    if 'cuda' in devices:
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA is not available, but input tensors are on CUDA.")
        cuda_devices = {t.device for t in tensors if t.device.type == 'cuda'}
        if len(cuda_devices) > 1:
            raise RuntimeError(f"Input tensors are on multiple CUDA devices: {cuda_devices}")
        return list(cuda_devices)[0]
    
    if torch.cuda.is_available():
        return torch.device('cuda')
    else:
        raise RuntimeError("Triton kernels require a CUDA-enabled GPU, but none was found.")


def run(*args, **kwargs):
    """
    Entry point for the GQA Ragged Prefill Causal Attention kernel.

    Args:
        q (torch.Tensor): Query tensor of shape [total_q, num_qo_heads, head_dim].
        k (torch.Tensor): Key tensor of shape [total_kv, num_kv_heads, head_dim].
        v (torch.Tensor): Value tensor of shape [total_kv, num_kv_heads, head_dim].
        qo_indptr (torch.Tensor): Query offsets for each sequence of shape [len_indptr].
        kv_indptr (torch.Tensor): Key-value offsets for each sequence of shape [len_indptr].
        sm_scale (float, optional): Softmax scale. Defaults to 1/sqrt(head_dim).

    Returns:
        Tuple[torch.Tensor, torch.Tensor]:
            - output (torch.Tensor): Attention output of shape [total_q, num_qo_heads, head_dim].
            - lse (torch.Tensor): Log-sum-exp of attention logits of shape [total_q, num_qo_heads].
    """
    # 1. Argument parsing
    arg_names = ['q', 'k', 'v', 'qo_indptr', 'kv_indptr', 'sm_scale']
    expected_arg_count = 5
    
    if len(args) > len(arg_names):
        raise TypeError(f"run() takes at most {len(arg_names)} positional arguments but {len(args)} were given")

    params = {name: val for name, val in zip(arg_names, args)}
    params.update(kwargs)

    missing_args = [name for name in arg_names[:expected_arg_count] if name not in params]
    if missing_args:
        raise TypeError(f"run() missing {len(missing_args)} required positional argument(s): {', '.join(missing_args)}")

    q, k, v, qo_indptr, kv_indptr = [params[name] for name in arg_names[:expected_arg_count]]
    sm_scale = params.get('sm_scale')

    # 2. Constants and shape assertions
    NUM_QO_HEADS = 32
    NUM_KV_HEADS = 4
    HEAD_DIM = 128
    
    total_q, num_qo_heads, head_dim = q.shape
    total_kv, num_kv_heads, _ = k.shape
    len_indptr = qo_indptr.shape[0]

    assert num_qo_heads == NUM_QO_HEADS, f"Expected num_qo_heads={NUM_QO_HEADS}, got {num_qo_heads}"
    assert num_kv_heads == NUM_KV_HEADS, f"Expected num_kv_heads={NUM_KV_HEADS}, got {num_kv_heads}"
    assert head_dim == HEAD_DIM, f"Expected head_dim={HEAD_DIM}, got {head_dim}"
    assert qo_indptr.dim() == 1 and kv_indptr.dim() == 1, "indptr tensors must be 1D"
    assert len_indptr > 0, "indptr tensors cannot be empty"
    assert total_q == qo_indptr[-1].item(), f"total_q ({total_q}) must match qo_indptr[-1] ({qo_indptr[-1].item()})"
    assert total_kv == kv_indptr[-1].item(), f"total_kv ({total_kv}) must match kv_indptr[-1] ({kv_indptr[-1].item()})"
    assert qo_indptr.shape == kv_indptr.shape, "qo_indptr and kv_indptr must have the same shape"

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

    # 3. Device management
    initial_device = q.device
    kernel_device = _get_device(q, k, v, qo_indptr, kv_indptr)
    
    q, k, v, qo_indptr, kv_indptr = [t.to(kernel_device) for t in [q, k, v, qo_indptr, kv_indptr]]
    
    q, k, v = [t.contiguous() for t in [q, k, v]]

    # 4. Prepare kernel inputs and outputs
    output = torch.empty_like(q, dtype=torch.bfloat16)
    lse = torch.full((total_q, NUM_QO_HEADS), -float("inf"), dtype=torch.float32, device=kernel_device)

    # 5. Launch kernel
    grid = (total_q, NUM_QO_HEADS)
    
    BLOCK_N = 64
    
    if total_q > 0:
        # Precompute a mapping from global query index to batch index for efficient lookup in the kernel
        q_indices = torch.arange(total_q, device=kernel_device)
        qo_ends = qo_indptr[1:]
        q_to_b_idx = torch.searchsorted(qo_ends, q_indices, right=True)
        
        gqa_ragged_prefill_causal_h32_kv4_d128_kernel[grid](
            q, k, v,
            qo_indptr, kv_indptr, q_to_b_idx,
            output, lse,
            sm_scale,
            q.stride(0), q.stride(1),
            k.stride(0), k.stride(1),
            v.stride(0), v.stride(1),
            total_q,
            GQA_RATIO=NUM_QO_HEADS // NUM_KV_HEADS,
            NUM_QO_HEADS=NUM_QO_HEADS,
            HEAD_DIM=HEAD_DIM,
            BLOCK_N=BLOCK_N,
        )

    # 6. Restore output device
    output = output.to(initial_device)
    lse = lse.to(initial_device)

    return output, lse
scrolls · 239 lines total

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

Best evidence level for this revision: reported

JSON