Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / triton0b9300

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-0b9300?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:b46215b392aba06da4c41e62b2c363cfdeec1ce27d516ff516d625fe5fe38555
license declaredApache-2.0
license concludedApache-2.0
authorsgemini-2.5-pro
imported2026-08-20

Kernel source

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

@triton.jit
def _top_k_sampling_kernel(
    probs_ptr,
    top_k_ptr,
    samples_ptr,
    rand_seed,
    BATCH_SIZE: tl.constexpr,
    VOCAB_SIZE: tl.constexpr,
    BLOCK_V: tl.constexpr,
):
    """
    Triton kernel for top-k sampling.

    This kernel performs top-k sampling for a batch of probability distributions.
    For each row, it filters the distribution to keep only the top `k` probabilities,
    renormalizes them, and then samples a single token using multinomial sampling.

    Note on implementation:
    A true top-k selection requires sorting or a complex parallel selection algorithm
    (like quickselect), which is hard to implement efficiently in a single Triton kernel
    for a large vocabulary and dynamic `k`. This implementation uses a highly efficient
    binary search method to find a probability threshold that approximates the k-th
    largest probability. This is technically a top-p (nucleus) sampling approach where `p`
    is chosen to correspond to `k` elements. This is a common high-performance strategy.
    It may differ from a strict index-based `torch.argsort` approach in cases of
    probabilities with identical values at the k-th position, but provides a massive
    performance boost over naive implementations.

    Grid: (BATCH_SIZE,)
    Each program in the grid handles one sequence in the batch.
    """
    # Program ID corresponds to the batch index
    pid = tl.program_id(0)

    # --- Step 1: Load `k` and determine if filtering is needed ---
    # `k` is specific to each sequence in the batch
    k = tl.load(top_k_ptr + pid)
    do_filter = (k > 0) & (k < VOCAB_SIZE)

    # Pointer to the start of the current row's probabilities
    row_probs_ptr = probs_ptr + pid * VOCAB_SIZE
    
    threshold = -1.0
    probs_sum = 1.0  # Default value, will be re-calculated if filtering occurs

    # --- Step 2: Top-k filtering logic ---
    if do_filter:
        # --- 2a: Find the threshold (approximating k-th largest value) via binary search ---
        min_p = 0.0
        
        # First, find the maximum probability in the row to establish a tight search range [0, max_p]
        max_p_val = tl.zeros((), dtype=tl.float32)
        v_offsets = tl.arange(0, BLOCK_V)
        for _ in range(0, tl.cdiv(VOCAB_SIZE, BLOCK_V)):
            mask = v_offsets < VOCAB_SIZE
            p = tl.load(row_probs_ptr + v_offsets, mask=mask, other=0.0)
            
            block_max = tl.max(p, axis=0)
            max_p_val = tl.maximum(max_p_val, block_max)
            
            v_offsets += BLOCK_V
        
        # Binary search for the threshold value. 16 iterations provide good precision for fp32.
        for _ in range(16):
            pivot = (min_p + max_p_val) * 0.5
            # Count how many probabilities are >= the pivot
            count = tl.zeros((), dtype=tl.int32)
            v_offsets = tl.arange(0, BLOCK_V)
            for _ in range(0, tl.cdiv(VOCAB_SIZE, BLOCK_V)):
                mask = v_offsets < VOCAB_SIZE
                p = tl.load(row_probs_ptr + v_offsets, mask=mask, other=0.0)
                count += tl.sum((p >= pivot).to(tl.int32))
                v_offsets += BLOCK_V
            
            # Adjust the search range based on the count
            if count >= k:
                min_p = pivot
            else:
                max_p_val = pivot
        
        threshold = min_p

        # --- 2b: Calculate the sum of the filtered probabilities for normalization ---
        current_sum = tl.zeros((), dtype=tl.float32)
        v_offsets = tl.arange(0, BLOCK_V)
        for _ in range(0, tl.cdiv(VOCAB_SIZE, BLOCK_V)):
            mask = v_offsets < VOCAB_SIZE
            p = tl.load(row_probs_ptr + v_offsets, mask=mask, other=0.0)
            p_filtered = tl.where(p >= threshold, p, 0.0)
            current_sum += tl.sum(p_filtered)
            v_offsets += BLOCK_V
        probs_sum = current_sum
    else:
        # If no filtering, compute sum of all probabilities for numerical stability
        current_sum = tl.zeros((), dtype=tl.float32)
        v_offsets = tl.arange(0, BLOCK_V)
        for _ in range(0, tl.cdiv(VOCAB_SIZE, BLOCK_V)):
            mask = v_offsets < VOCAB_SIZE
            p = tl.load(row_probs_ptr + v_offsets, mask=mask, other=0.0)
            current_sum += tl.sum(p)
            v_offsets += BLOCK_V
        probs_sum = current_sum

    # --- Step 3: Multinomial Sampling ---
    # Use philox for high-quality pseudo-random numbers.
    philox_offset = pid.to(tl.uint64)
    rand_val = tl.rand(rand_seed, philox_offset)
    
    # Scale the random number by the sum of probabilities to get the target for the cumulative sum
    target_cumulative_prob = rand_val * probs_sum

    # Scan through the distribution to find the token corresponding to the random sample.
    cumulative_prob = tl.zeros((), dtype=tl.float32)
    # Initialize result index to a large value to act as a sentinel.
    final_idx = tl.full((), VOCAB_SIZE * 2, dtype=tl.int64)
    
    v_offsets = tl.arange(0, BLOCK_V)
    for _ in range(0, tl.cdiv(VOCAB_SIZE, BLOCK_V)):
        mask = v_offsets < VOCAB_SIZE
        probs = tl.load(row_probs_ptr + v_offsets, mask=mask, other=0.0)
        indices = v_offsets.to(tl.int64)

        if do_filter:
            probs = tl.where(probs >= threshold, probs, 0.0)
        
        # Calculate cumulative sum within the block and add the sum from previous blocks
        block_cumsum = tl.cumsum(probs, axis=0)
        total_cumsum = cumulative_prob + block_cumsum
        
        # Identify candidates: indices where the cumulative sum crosses the target
        is_candidate = (total_cumsum > target_cumulative_prob)
        
        # Check if this block is the first to contain candidates
        is_winning_block = cumulative_prob <= target_cumulative_prob

        if is_winning_block:
            candidate_indices = tl.where(is_candidate, indices, final_idx)
            # The minimum of these is the first valid index in this block
            # FIX: The original `tl.reduce(..., tl.min)` caused a CompilationError.
            # The idiomatic and correct way to perform this reduction is to use
            # `tl.min(tensor, axis=0)`.
            block_min_idx = tl.min(candidate_indices, axis=0)
            # Update the overall final index with the minimum found so far
            final_idx = tl.minimum(final_idx, block_min_idx)

        cumulative_prob += tl.sum(probs)
        v_offsets += BLOCK_V
        
    # --- Step 4: Finalize and store the result ---
    # Handle edge cases where sum of probabilities is zero or rounding errors occur.
    final_idx = tl.where(probs_sum > 0.0, final_idx, 0)
    final_idx = tl.where(final_idx < VOCAB_SIZE, final_idx, VOCAB_SIZE - 1)
        
    tl.store(samples_ptr + pid, final_idx)


def run(*args, **kwargs):
    """
    Wrapper function for the top-k sampling Triton kernel.

    This function handles device management, kernel launching, and tensor validation.
    It ensures that input tensors are on the correct GPU device and that the output
    tensor is moved back to the original device of the input tensors.

    Args:
        *args: Positional arguments. Expects `probs` and `top_k`.
        **kwargs: Keyword arguments. Expects `probs` and `top_k`.

    Returns:
        torch.Tensor: A tensor of shape [batch_size] containing the sampled token indices.
    """
    # --- Argument Parsing and Validation ---
    if args and kwargs:
        raise ValueError("Cannot provide both positional and keyword arguments.")
    
    if args:
        if len(args) != 2:
            raise ValueError(f"Expected 2 positional arguments (`probs`, `top_k`), but got {len(args)}.")
        probs, top_k = args
    elif kwargs:
        if "probs" not in kwargs or "top_k" not in kwargs:
            raise ValueError("Missing required keyword arguments: `probs` and `top_k`.")
        probs = kwargs.pop("probs")
        top_k = kwargs.pop("top_k")
        if kwargs:
            raise ValueError(f"Unexpected keyword arguments: {list(kwargs.keys())}")
    else:
        raise ValueError("No arguments provided. Expected `probs` and `top_k`.")

    # --- Shape and DType Validation ---
    if not isinstance(probs, torch.Tensor):
        raise TypeError(f"`probs` must be a torch.Tensor, but got {type(probs)}")
    if not isinstance(top_k, torch.Tensor):
        raise TypeError(f"`top_k` must be a torch.Tensor, but got {type(top_k)}")

    if probs.ndim != 2:
        raise ValueError(f"Expected `probs` to be a 2D tensor, but got {probs.ndim} dimensions.")
    
    batch_size, vocab_size = probs.shape
    VOCAB_SIZE = 151936
    if vocab_size != VOCAB_SIZE:
        raise ValueError(f"Expected `probs` to have vocab_size={VOCAB_SIZE}, but got {vocab_size}.")

    if top_k.ndim != 1 or top_k.shape[0] != batch_size:
        raise ValueError(f"Expected `top_k` to be a 1D tensor of size {batch_size}, but got shape {top_k.shape}.")

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

    device = probs.device
    original_device_str = 'cpu' if device.type == 'cpu' else device.type
    
    gpu_device = 'cuda' # Assume we run on the default CUDA device
    probs_gpu = probs.to(gpu_device, non_blocking=True)
    top_k_gpu = top_k.to(gpu_device, non_blocking=True)

    # --- Kernel Launch ---
    # Ensure contiguous tensors for performance
    probs_gpu = probs_gpu.contiguous().to(torch.float32)
    top_k_gpu = top_k_gpu.contiguous().to(torch.int32)
    
    # Create output tensor
    samples = torch.empty(batch_size, dtype=torch.int64, device=gpu_device)
    
    # Create a random seed for the kernel
    rand_seed = torch.randint(0, 2**63 - 1, (1,), dtype=torch.int64, device='cpu').item()

    # Configure grid and block size
    grid = (batch_size,)
    BLOCK_V = 2048
    
    _top_k_sampling_kernel[grid](
        probs_ptr=probs_gpu,
        top_k_ptr=top_k_gpu,
        samples_ptr=samples,
        rand_seed=rand_seed,
        BATCH_SIZE=batch_size,
        VOCAB_SIZE=VOCAB_SIZE,
        BLOCK_V=BLOCK_V
    )
    
    # --- Output Device Management ---
    # Move result back to the original device of the inputs
    if original_device_str != 'cuda':
        samples = samples.to(device, non_blocking=True)
    
    return samples
scrolls · 254 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON