Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / tritonrbz3hy

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-rbz3hy?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:dbbee34fd020694337a0db4c42071daf911f49664f91fdb56a170471ade15444
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(
mmas = tl.dot(q_mat, tl.trans(k)) * sm_scale
num-warps = 4triton.Config({'BLOCK_N': 64}, num_warps=4, num_stages=3),
online-softmaxm_i_new = tl.maximum(m_i, tl.max(s, axis=1))
stages = 3triton.Config({'BLOCK_N': 64}, num_warps=4, num_stages=3),

Kernel source

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

# `triton.jit`'ed functions can be auto-tuned by using the `triton.autotune` decorator, which consumes:
#   - A list of `triton.Config` objects that define different configurations of values for user-defined arguments.
#   - A `key` argument containing a list of names of arguments used to determine which configuration is chosen.
@triton.autotune(
    configs=[
        triton.Config({'BLOCK_N': 64}, num_warps=4, num_stages=3),
        triton.Config({'BLOCK_N': 128}, num_warps=4, num_stages=3),
        triton.Config({'BLOCK_N': 256}, num_warps=8, num_stages=2),
        triton.Config({'BLOCK_N': 128}, num_warps=8, num_stages=4),
        triton.Config({'BLOCK_N': 64}, num_warps=4, num_stages=4),
        triton.Config({'BLOCK_N': 32}, num_warps=2, num_stages=2),
    ],
    key=['HEAD_DIM'],
)
@triton.jit
def gqa_ragged_prefill_causal_kernel(
    # Pointers to matrices
    Q, K, V, O, LSE,
    # Pointer to precomputed location map
    q_loc,
    sm_scale,
    # Strides
    Q_stride_t, Q_stride_h, Q_stride_d,
    K_stride_t, K_stride_h, K_stride_d,
    V_stride_t, V_stride_h, V_stride_d,
    O_stride_t, O_stride_h, O_stride_d,
    LSE_stride_t, LSE_stride_h,
    q_loc_stride_t, q_loc_stride_d,
    # Compile-time constants
    NUM_QO_HEADS: tl.constexpr,
    NUM_KV_HEADS: tl.constexpr,
    HEAD_DIM: tl.constexpr,
    BLOCK_D: tl.constexpr,
    BLOCK_N: tl.constexpr,
):
    """
    Triton kernel for Grouped-Query Attention for ragged prefill with causal masking.

    This kernel computes attention for one query token against its corresponding
    key-value sequence. The grid is launched with one program per (query_token, query_head).

    The ragged nature of the input is handled by a precomputed location map `q_loc`,
    which provides sequence boundaries for each query token, avoiding complex and slow
    indexing logic inside the kernel.

    The computation uses a tiled approach similar to FlashAttention to efficiently
    process the key-value sequence in blocks, leveraging shared memory implicitly
    via Triton's dot product operations. Online softmax is used to maintain
    numerical stability and compute the result in a single pass over the KV cache.
    """
    # Grid is (total_q, num_qo_heads)
    pid_qt = tl.program_id(0)  # Global query token index
    pid_h = tl.program_id(1)   # Query head index

    # --- 1. Load sequence boundaries and determine context ---
    # Load [q_start, q_end, kv_start, kv_end] from the precomputed map
    q_loc_ptr = q_loc + pid_qt * q_loc_stride_t
    q_start = tl.load(q_loc_ptr + 0 * q_loc_stride_d)
    q_end = tl.load(q_loc_ptr + 1 * q_loc_stride_d)
    kv_start = tl.load(q_loc_ptr + 2 * q_loc_stride_d)
    kv_end = tl.load(q_loc_ptr + 3 * q_loc_stride_d)

    # Calculate local query index and sequence lengths
    q_idx_local = pid_qt - q_start
    num_q_tokens = q_end - q_start
    num_kv_tokens = kv_end - kv_start
    delta = num_kv_tokens - num_q_tokens
    
    # Causal sequence length for this query
    kv_len_for_q = tl.minimum(q_idx_local + 1 + delta, num_kv_tokens)

    # --- 2. Determine head indices and pointers ---
    GQA_RATIO: tl.constexpr = NUM_QO_HEADS // NUM_KV_HEADS
    kv_head_idx = pid_h // GQA_RATIO

    # Pointers to K and V for the correct sequence and head
    k_batch_head_ptr = K + kv_start * K_stride_t + kv_head_idx * K_stride_h
    v_batch_head_ptr = V + kv_start * V_stride_t + kv_head_idx * V_stride_h

    # --- 3. Initialize accumulator and online softmax statistics ---
    acc = tl.zeros([BLOCK_D], dtype=tl.float32)
    m_i = -float('inf')
    l_i = 0.0

    # --- 4. Load query vector ---
    q_ptr = Q + pid_qt * Q_stride_t + pid_h * Q_stride_h
    offs_d = tl.arange(0, BLOCK_D)
    q = tl.load(q_ptr + offs_d, mask=offs_d < HEAD_DIM, other=0.0).to(tl.float32)

    # --- 5. Main loop over KV sequence blocks ---
    kv_offset = 0
    # The loop condition handles cases where kv_len_for_q <= 0
    while kv_offset < kv_len_for_q:
        # Pointers to the current block of K and V
        k_ptr = k_batch_head_ptr + kv_offset * K_stride_t
        v_ptr = v_batch_head_ptr + kv_offset * V_stride_t

        # Offsets for loading K and V blocks
        offs_n = tl.arange(0, BLOCK_N)
        k_offs = (offs_n[:, None] * K_stride_t + offs_d[None, :])
        v_offs = (offs_n[:, None] * V_stride_t + offs_d[None, :])

        # Create a mask for the current block to handle both padding within the
        # block and the causal boundary.
        k_mask = (kv_offset + offs_n) < kv_len_for_q

        # Load K and V blocks with masking
        k = tl.load(k_ptr + k_offs, mask=k_mask[:, None] & (offs_d[None, :] < HEAD_DIM), other=0.0).to(tl.float32)
        v = tl.load(v_ptr + v_offs, mask=k_mask[:, None] & (offs_d[None, :] < HEAD_DIM), other=0.0)

        # --- Core attention computation (FIXED BLOCK) ---
        # Compute Q @ K.T
        # FIX: The core issue was that tl.dot requires 2D inputs. We reshape the
        # 1D query vector `q` into a 2D matrix `q_mat` of shape [1, BLOCK_D].
        q_mat = tl.reshape(q, (1, BLOCK_D))
        s = tl.dot(q_mat, tl.trans(k)) * sm_scale
        
        # Apply mask to logits. s is [1, N], k_mask is [N], broadcasting is fine.
        s = tl.where(k_mask, s, -float('inf'))

        # --- Online softmax update ---
        # FIX: Since `s` is now 2D [1, N], the reductions must be handled correctly.
        # We reduce over axis=1 to get [1]-shaped tensors for the statistics.
        m_i_new = tl.maximum(m_i, tl.max(s, axis=1))
        alpha = tl.exp(m_i - m_i_new)
        p = tl.exp(s - m_i_new)
        l_i_new = alpha * l_i + tl.sum(p, axis=1)

        # Update accumulator
        # Triton correctly broadcasts the [1]-shaped `alpha` tensor across `acc`.
        acc = acc * alpha
        # FIX: `p` is [1, N], `v` is [N, D]. Dot product gives [1, D].
        # We must reshape the result to [D] to correctly add it to `acc`.
        delta_acc = tl.dot(p.to(v.dtype), v)
        acc += tl.reshape(delta_acc, (BLOCK_D,))

        # Update statistics for next iteration
        m_i = m_i_new
        l_i = l_i_new

        kv_offset += BLOCK_N

    # --- 6. Finalize and store results ---
    # Finalize accumulator
    l_i_safe = tl.where(l_i == 0.0, 1.0, l_i)
    acc = acc / l_i_safe

    # Compute 2-based log-sum-exp
    log2_e = 1.4426950408889634  # 1.0 / ln(2)
    lse = m_i + tl.log(l_i)
    lse = lse * log2_e
    # If all scores were -inf, l_i is 0, log(l_i) is -inf, which is correct.
    lse = tl.where(l_i == 0.0, -float('inf'), lse)

    # Store output and LSE
    offs_d_store = tl.arange(0, BLOCK_D)
    o_ptr = O + pid_qt * O_stride_t + pid_h * O_stride_h
    lse_ptr = LSE + pid_qt * LSE_stride_t + pid_h * LSE_stride_h

    tl.store(o_ptr + offs_d_store, acc.to(O.dtype.element_ty), mask=offs_d_store < HEAD_DIM)
    tl.store(lse_ptr, lse)


def gqa_ragged_prefill_causal_h32_kv8_d128(q, k, v, qo_indptr, kv_indptr, sm_scale):
    """
    Wrapper function for the GQA ragged prefill kernel.

    This function prepares tensors, defines the launch grid, and calls the
    Triton kernel. It includes a host-side precomputation step to create
    a `q_loc` map, which massively simplifies the kernel's indexing logic
    by providing each query token with its sequence boundaries directly.
    """
    # Extract shape information
    total_q, num_qo_heads, head_dim = q.shape
    total_kv, num_kv_heads, _ = k.shape
    len_indptr = qo_indptr.shape[0]
    batch_size = len_indptr - 1

    # Allocate output tensors
    output = torch.empty_like(q)
    lse = torch.empty((total_q, num_qo_heads), dtype=torch.float32, device=q.device)

    # Precompute location map: [q_start, q_end, kv_start, kv_end] for each query
    # This avoids a complex and slow search/lookup within the kernel.
    q_loc = torch.empty((total_q, 4), dtype=torch.int32, device=q.device)
    if batch_size > 0 and total_q > 0:
        for b in range(batch_size):
            q_s, q_e = qo_indptr[b].item(), qo_indptr[b+1].item()
            kv_s, kv_e = kv_indptr[b].item(), kv_indptr[b+1].item()
            if q_s < q_e:
                # Use broadcasting to fill the map for all tokens in the sequence
                q_loc[q_s:q_e, 0] = q_s
                q_loc[q_s:q_e, 1] = q_e
                q_loc[q_s:q_e, 2] = kv_s
                q_loc[q_s:q_e, 3] = kv_e

    # Define the launch grid: one program per (query_token, query_head)
    grid = (total_q, num_qo_heads)

    # Call the Triton kernel, only if there are tokens to process
    if total_q > 0:
        gqa_ragged_prefill_causal_kernel[grid](
            q, k, v, output, lse,
            q_loc,
            sm_scale,
            # Strides
            q.stride(0), q.stride(1), q.stride(2),
            k.stride(0), k.stride(1), k.stride(2),
            v.stride(0), v.stride(1), v.stride(2),
            output.stride(0), output.stride(1), output.stride(2),
            lse.stride(0), lse.stride(1),
            q_loc.stride(0), q_loc.stride(1),
            # Constants
            NUM_QO_HEADS=num_qo_heads,
            NUM_KV_HEADS=num_kv_heads,
            HEAD_DIM=head_dim,
            BLOCK_D=head_dim,
            # BLOCK_N is autotuned
        )

    return output, lse


def run(*args, **kwargs):
    """
    Public entry point for the operation.

    Handles device management, argument parsing, and calls the underlying
    Triton implementation. It ensures that input tensors are on the correct
    device (CUDA) and that output tensors are moved back to the original
    device.
    """
    # --- Argument Parsing ---
    # This robustly handles both positional and keyword arguments.
    arg_names = ['q', 'k', 'v', 'qo_indptr', 'kv_indptr', 'sm_scale']
    arg_dict = {name: kwargs.get(name) for name in arg_names}
    for i, arg in enumerate(args):
        # This will overwrite a kwarg if it was also passed as an arg,
        # which is standard Python behavior.
        if i < len(arg_names):
             arg_dict[arg_names[i]] = arg

    q = arg_dict['q']
    k = arg_dict['k']
    v = arg_dict['v']
    qo_indptr = arg_dict['qo_indptr']
    kv_indptr = arg_dict['kv_indptr']
    sm_scale = arg_dict['sm_scale']

    # Check for missing required arguments
    required_args = ['q', 'k', 'v', 'qo_indptr', 'kv_indptr']
    for arg_name in required_args:
        if arg_dict[arg_name] is None:
            raise TypeError(f"Missing required argument: '{arg_name}'")


    # --- Constants and Defaults ---
    HEAD_DIM = 128
    if sm_scale is None:
        sm_scale = 1.0 / math.sqrt(HEAD_DIM)

    # --- Device Management ---
    if not torch.cuda.is_available():
        raise RuntimeError("This Triton kernel requires a CUDA-enabled GPU.")

    original_device = q.device
    is_cpu = original_device.type == 'cpu'

    if is_cpu:
        # Move all tensors to GPU
        q, k, v, qo_indptr, kv_indptr = (
            t.cuda() for t in [q, k, v, qo_indptr, kv_indptr]
        )
    elif q.device.type != 'cuda':
        raise RuntimeError(f"Unsupported device: {q.device}. Only CPU and CUDA are supported.")

    # --- Constraints Validation ---
    total_q, num_qo_heads, head_dim = q.shape
    _, num_kv_heads, _ = k.shape

    assert num_qo_heads == 32, f"Expected num_qo_heads=32, but got {num_qo_heads}"
    assert num_kv_heads == 8, f"Expected num_kv_heads=8, but got {num_kv_heads}"
    assert head_dim == HEAD_DIM, f"Expected head_dim={HEAD_DIM}, but got {head_dim}"
    assert total_q == qo_indptr[-1].item(), "total_q must match qo_indptr[-1]"
    assert k.shape[0] == kv_indptr[-1].item(), "total_kv must match kv_indptr[-1]"

    # --- Kernel Execution ---
    output, lse = gqa_ragged_prefill_causal_h32_kv8_d128(q, k, v, qo_indptr, kv_indptr, sm_scale)

    # --- Device Restoration ---
    if is_cpu:
        output = output.to(original_device)
        lse = lse.to(original_device)

    return output, lse
scrolls · 300 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON