Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / tritonf8ce0a

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-f8ce0a?include=source"
interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp32, 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:e3706a7eb0b412bf452df3285c737486b0efce400e39f9cb6002a57765d900f6
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_V': 1024, 'MAX_K_BUFFER': 2048}, num_warps=4),

Kernel source

main.py264 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
#     meta-parameters (e.g., `BLOCK_SIZE_M`) and compiler options (e.g., `num_warps`)
#   - A `key` argument containing argument names for the kernel parameters
#
# JITed functions can be decorated with `triton.autotune` to optimize for a given input shape.
# This is especially important for kernels that handle tensors with variable shapes.
# The `key` argument is used to lookup the best configuration for a given set of input shapes.
# For this kernel, we tune for batch_size.
@triton.autotune(
    configs=[
        triton.Config({'BLOCK_V': 1024, 'MAX_K_BUFFER': 2048}, num_warps=4),
        triton.Config({'BLOCK_V': 2048, 'MAX_K_BUFFER': 2048}, num_warps=8),
        triton.Config({'BLOCK_V': 1024, 'MAX_K_BUFFER': 4096}, num_warps=4),
        triton.Config({'BLOCK_V': 2048, 'MAX_K_BUFFER': 4096}, num_warps=8),
        triton.Config({'BLOCK_V': 512, 'MAX_K_BUFFER': 2048}, num_warps=2),
    ],
    key=['batch_size'],
)
@triton.jit
def top_k_top_p_sampling_from_probs_v151936_kernel(
    # Pointers to tensors
    probs_ptr,
    top_k_ptr,
    top_p_ptr,
    samples_ptr,
    # Random sampling state
    seed,
    offsets_ptr,
    # Tensor dimensions
    batch_size,
    # Strides
    stride_probs_b,
    # Meta-parameters
    VOCAB_SIZE: tl.constexpr,
    BLOCK_V: tl.constexpr,
    MAX_K_BUFFER: tl.constexpr,
):
    """
    Triton kernel for top-k/top-p sampling.
    Each program instance processes one sequence from the batch.
    """
    # -----------------------------------------------------------
    # Program setup
    # -----------------------------------------------------------
    pid = tl.program_id(0)

    # Load per-sequence sampling parameters
    k = tl.load(top_k_ptr + pid)
    p = tl.load(top_p_ptr + pid)

    # Pointers to the current sequence's data
    probs_row_ptr = probs_ptr + pid * stride_probs_b

    # -----------------------------------------------------------
    # Greedy decoding path (p <= 0.0) -> argmax
    # -----------------------------------------------------------
    if p <= 0.0:
        max_prob = -1.0
        max_idx = -1
        v_offsets = tl.arange(0, BLOCK_V)
        for v_start in range(0, VOCAB_SIZE, BLOCK_V):
            v_range = v_start + v_offsets
            v_mask = v_range < VOCAB_SIZE
            row_probs = tl.load(probs_row_ptr + v_range, mask=v_mask, other=-1.0)

            block_max_prob = tl.max(row_probs)
            
            # If the max in this block is greater than the global max, update global max
            # and find the first index of this new max in the current block.
            if block_max_prob > max_prob:
                max_prob = block_max_prob
                is_max = (row_probs == max_prob) & v_mask
                max_indices_in_block = tl.where(is_max, v_range, VOCAB_SIZE + 1)
                max_idx = tl.min(max_indices_in_block)
            # If the block max is equal to the global max, we only update the index
            # if the new index is smaller (torch.argmax behavior).
            elif block_max_prob == max_prob:
                is_max = (row_probs == max_prob) & v_mask
                max_indices_in_block = tl.where(is_max, v_range, VOCAB_SIZE + 1)
                block_min_idx = tl.min(max_indices_in_block)
                if block_min_idx < max_idx:
                    max_idx = block_min_idx

        tl.store(samples_ptr + pid, max_idx)
        return

    # -----------------------------------------------------------
    # Top-K and Top-P Sampling Path
    # -----------------------------------------------------------

    # --- Stage 1: Find top candidates using a streaming approach ---
    # `effective_k` is the number of candidates to consider after sorting.
    # If k is invalid (<=0) or too large, we default to the buffer size for candidate search,
    # but k will be respected during filtering.
    effective_k = k
    if k <= 0 or k > MAX_K_BUFFER:
        effective_k = MAX_K_BUFFER

    # Initialize SRAM buffers with the first block of candidates.
    v_offsets_init = tl.arange(0, MAX_K_BUFFER)
    v_mask_init = v_offsets_init < VOCAB_SIZE
    sram_probs = tl.load(probs_row_ptr + v_offsets_init, mask=v_mask_init, other=-1.0)
    sram_indices = v_offsets_init.to(tl.int32)

    min_prob_in_sram = tl.min(sram_probs)

    # Iterate over the rest of the vocabulary to find better candidates.
    v_offsets = tl.arange(0, BLOCK_V)
    for v_start in range(MAX_K_BUFFER, VOCAB_SIZE, BLOCK_V):
        v_range = v_start + v_offsets
        v_mask = v_range < VOCAB_SIZE
        block_probs = tl.load(probs_row_ptr + v_range, mask=v_mask, other=-1.0)

        # Optimization: only process block if it contains a potential candidate
        if tl.max(block_probs) > min_prob_in_sram:
            # This loop is unrolled by the compiler. It serially updates the candidate set.
            for i in range(BLOCK_V):
                prob = tl.load(probs_row_ptr + v_start + i, mask=(v_start + i < VOCAB_SIZE), other=-1.0)
                if prob > min_prob_in_sram:
                    # Find the location of the minimum element and replace it.
                    min_mask = sram_probs == min_prob_in_sram
                    # To break ties, take the one with the smallest index in the sram buffer.
                    min_indices = tl.where(min_mask, tl.arange(0, MAX_K_BUFFER), MAX_K_BUFFER + 1)
                    first_min_idx = tl.min(min_indices)

                    # Replace the minimum element with the new, larger candidate.
                    sram_probs = tl.where(tl.arange(0, MAX_K_BUFFER) == first_min_idx, prob, sram_probs)
                    sram_indices = tl.where(tl.arange(0, MAX_K_BUFFER) == first_min_idx, v_start + i, sram_indices)

                    # Update the minimum for the next iteration of this inner loop.
                    min_prob_in_sram = tl.min(sram_probs)

    # Sort the final candidates before top-p filtering.
    # We use a robust packing method to sort key-value pairs.
    packed = (sram_probs * 2147483647.0).to(tl.int32).to(tl.int64) << 32 | sram_indices.to(tl.int64)
    sorted_packed = tl.sort(packed, descending=True)
    
    # Corrected unpacking: Do not mask the sign bit.
    sram_probs = (sorted_packed >> 32).to(tl.int32).to(tl.float32) / 2147483647.0
    sram_indices = (sorted_packed & 0xFFFFFFFF).to(tl.int32)

    # --- Stage 2: Apply Top-K then Top-P filtering on the candidates ---
    k_arange = tl.arange(0, MAX_K_BUFFER)
    k_mask = k_arange < effective_k
    masked_sram_probs = tl.where(k_mask, sram_probs, 0.0)

    total_prob_k = tl.sum(masked_sram_probs, axis=0)
    norm_sram_probs = masked_sram_probs / (total_prob_k + 1e-9)

    cumsum_probs = tl.cumsum(norm_sram_probs, axis=0)

    # Condition for discarding token `i` is when cumulative prob of tokens `0..i-1` >= `p`.
    p_mask = (cumsum_probs - norm_sram_probs) >= p
    p_cutoff_idx_raw = tl.where(p_mask, k_arange, MAX_K_BUFFER)
    num_final_candidates = tl.min(p_cutoff_idx_raw)

    # Ensure at least one token is considered, matching reference behavior.
    if num_final_candidates == 0:
        num_final_candidates = 1
    # The number of final candidates cannot exceed the top-k limit.
    if num_final_candidates > effective_k:
        num_final_candidates = effective_k

    final_mask = k_arange < num_final_candidates
    final_probs = tl.where(final_mask, norm_sram_probs, 0.0)
    total_prob_p = tl.sum(final_probs, axis=0)

    # --- Stage 3: Sample from the final candidates ---
    rand_offset = tl.load(offsets_ptr + pid)
    tl.store(offsets_ptr + pid, rand_offset + 1)

    random_uniform = tl.rand(seed, rand_offset)
    random_scaled = random_uniform * (total_prob_p + 1e-9)

    final_cumsum = tl.cumsum(final_probs, axis=0)

    sampled_mask = random_scaled < final_cumsum
    sampled_idx_in_sram_raw = tl.where(sampled_mask, k_arange, MAX_K_BUFFER)
    sampled_idx_in_sram = tl.min(sampled_idx_in_sram_raw)

    # Gather the final token index from the sram buffer.
    selection_mask = k_arange == sampled_idx_in_sram
    final_sample_idx = tl.sum(tl.where(selection_mask, sram_indices, 0))

    tl.store(samples_ptr + pid, final_sample_idx.to(tl.int64))


def run(probs: torch.Tensor, top_k: torch.Tensor, top_p: torch.Tensor, **kwargs):
    """
    Wrapper function for the top_k_top_p_sampling Triton kernel.

    Args:
        probs (torch.Tensor): Probability distributions [batch_size, vocab_size], DType.FLOAT32.
        top_k (torch.Tensor): Number of top tokens to consider [batch_size], DType.INT32.
        top_p (torch.Tensor): Cumulative probability threshold [batch_size], DType.FLOAT32.

    Returns:
        torch.Tensor: Sampled token indices [batch_size], DType.INT64.
    """
    # -----------------------------------------------------------
    # Device and DType management
    # -----------------------------------------------------------
    original_device = probs.device

    if not torch.cuda.is_available():
        if any(t.is_cuda for t in [probs, top_k, top_p]):
            raise RuntimeError("CUDA is required for this Triton kernel, but is not available.")
        # This path is for CPU-only environments, which Triton doesn't support.
        raise RuntimeError("CUDA is required for this Triton kernel.")
    
    compute_device = torch.device('cuda')

    # Move all tensors to the compute device
    probs = probs.to(compute_device, non_blocking=True)
    top_k = top_k.to(compute_device, non_blocking=True)
    top_p = top_p.to(compute_device, non_blocking=True)

    # Ensure correct dtypes
    probs = probs.to(torch.float32)
    top_k = top_k.to(torch.int32)
    top_p = top_p.to(torch.float32)

    # -----------------------------------------------------------
    # Kernel launch setup
    # -----------------------------------------------------------
    batch_size, vocab_size = probs.shape

    if vocab_size != 151936:
        raise ValueError(f"This kernel is specialized for vocab_size=151936, but got {vocab_size}")

    samples = torch.empty(batch_size, dtype=torch.int64, device=compute_device)

    # Seed and offsets for random number generation
    seed = 1234
    offsets = torch.randint(0, vocab_size * 2, (batch_size,), dtype=torch.int32, device=compute_device)

    grid = lambda meta: (batch_size,)

    # -----------------------------------------------------------
    # Kernel invocation
    # -----------------------------------------------------------
    top_k_top_p_sampling_from_probs_v151936_kernel[grid](
        probs,
        top_k,
        top_p,
        samples,
        seed,
        offsets,
        batch_size,
        probs.stride(0),
        VOCAB_SIZE=vocab_size
    )

    # -----------------------------------------------------------
    # Finalization
    # -----------------------------------------------------------
    # Move the result back to the original device
    return samples.to(original_device)
scrolls · 264 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON