Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / triton8833c7

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Kernel source

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

@triton.jit
def top_k_top_p_sampling_from_probs_v129280_kernel(
    probs_ptr,
    top_k_ptr,
    top_p_ptr,
    samples_ptr,
    rand_seed,
    # VOCAB_SIZE is a constant, passed as a constexpr
    VOCAB_SIZE: tl.constexpr,
    # VOCAB_SIZE_P2 is the next power of 2 for sorting
    VOCAB_SIZE_P2: tl.constexpr,
):
    """
    Triton kernel for top-k/top-p sampling.

    This kernel processes one probability distribution per program instance.
    It performs a full sort of the vocabulary probabilities, which is memory-intensive.
    This implementation is chosen for its logical clarity and to fit within a single
    kernel. For very large vocabularies where the data exceeds L1/Shared Memory size,
    performance will be limited by memory spilling.

    For B200/Hopper architectures, the large register file and L2 cache can mitigate
    some spilling effects, but the algorithm remains memory-bound during the sort.

    Grid: (batch_size,)
    """
    # Get the batch index for this program instance
    batch_idx = tl.program_id(0)

    # --- 1. Load Data ---
    # Load top_k and top_p for the current sequence
    k = tl.load(top_k_ptr + batch_idx)
    p = tl.load(top_p_ptr + batch_idx)

    # Create pointers and ranges for the full vocabulary
    vocab_offsets = tl.arange(0, VOCAB_SIZE_P2)
    vocab_mask = vocab_offsets < VOCAB_SIZE
    
    # Load probabilities for the current sequence, padding with a negative value for the sort
    probs_row_ptr = probs_ptr + batch_idx * VOCAB_SIZE
    probs_vec = tl.load(probs_row_ptr + vocab_offsets, mask=vocab_mask, other=-1.0)
    
    # Create original indices, also padded
    indices_vec = tl.arange(0, VOCAB_SIZE_P2)

    # --- 2. Pack and Sort ---
    # FIX: The target Triton version does not support key-value sorting via `tl.sort((keys, values))`
    # and also lacks `tl.bitcast`. To work around this, we implement a manual packing scheme to
    # sort keys (probabilities) and values (indices) together. We scale the float32 probability
    # into the high bits of an int64 and place the int32 index into the low bits.

    # VOCAB_SIZE_P2 is 2^18, so indices need 18 bits.
    INDEX_BITS: tl.constexpr = 18
    # Use float64 for precision of the scaling factor.
    PROB_SCALE_FACTOR = (2.0 ** (63 - INDEX_BITS))

    # Scale probabilities and cast to int64. The order is preserved for positive values.
    # Negative probabilities (from padding) will correctly sort to the end.
    scaled_probs = (probs_vec * PROB_SCALE_FACTOR).to(tl.int64)

    # Combine scaled probabilities (high bits) and indices (low bits) into a single int64.
    packed_data = (scaled_probs << INDEX_BITS) + indices_vec.to(tl.int64)

    # Sort the packed data. Since probs are in the high bits, this sorts by probability.
    sorted_packed = tl.sort(packed_data, descending=True)

    # Unpack the sorted indices from the low bits of the sorted packed data.
    INDEX_MASK: tl.constexpr = (1 << INDEX_BITS) - 1
    sorted_indices = (sorted_packed & INDEX_MASK).to(tl.int32)

    # Re-gather the true probabilities using the sorted indices. This is a necessary
    # step because we only stored a scaled approximation in the packed data.
    gather_mask = sorted_indices < VOCAB_SIZE
    sorted_probs = tl.load(probs_row_ptr + sorted_indices, mask=gather_mask, other=0.0)
    
    # --- 3. Top-K Filtering ---
    # Determine the effective K. A k of 0 or >= vocab_size means no top-k filtering.
    use_top_k = (k > 0) & (k < VOCAB_SIZE)
    effective_k = tl.where(use_top_k, k, VOCAB_SIZE)

    # Create a mask for the top-k elements. Since `sorted_probs` is already zero-padded
    # for invalid indices from the gather step, we no longer need a separate validity mask here.
    k_arange = tl.arange(0, VOCAB_SIZE_P2)
    k_mask = k_arange < effective_k

    # Apply the mask to the sorted probabilities
    probs_after_k = tl.where(k_mask, sorted_probs, 0.0)
    
    # Renormalize the probabilities after top-k
    sum_probs_k = tl.sum(probs_after_k, axis=0)
    probs_after_k = probs_after_k / (sum_probs_k + 1e-9)

    # --- 4. Top-P (Nucleus) Filtering ---
    # This filtering is applied on the result of the top-k filtering.
    
    # FIX: Direct indexing like `sorted_indices[0]` is not supported on a tl.tensor.
    # We use a reduction with a mask to extract the first element as a scalar for the greedy sample.
    is_first_element_mask = k_arange == 0
    greedy_sample_tensor = tl.where(is_first_element_mask, sorted_indices, 0)
    greedy_sample = tl.sum(greedy_sample_tensor, axis=0)
    
    # Probabilities are already sorted, so we can compute the cumulative distribution
    cdf = tl.cumsum(probs_after_k, axis=0)
    
    # Find tokens to keep. A token is kept if its cumulative probability *before*
    # including itself is less than p.
    shifted_cdf = cdf - probs_after_k
    p_mask = (shifted_cdf < p) & k_mask

    # Apply the p_mask to the k-filtered probabilities
    probs_after_p = tl.where(p_mask, probs_after_k, 0.0)
    
    # Renormalize the probabilities after top-p
    sum_probs_p = tl.sum(probs_after_p, axis=0)
    probs_after_p = probs_after_p / (sum_probs_p + 1e-9)

    # --- 5. Sampling ---
    # Choose which distribution to sample from based on p
    # If p >= 1.0, top-p is a no-op, so we use the top-k filtered distribution.
    # If 0 < p < 1.0, use the top-p filtered distribution.
    final_probs = tl.where((p > 0.0) & (p < 1.0), probs_after_p, probs_after_k)

    # Generate a random number for this sequence
    # FIX: The prime number literal 2654435761 exceeds the int32 maximum.
    # Cast batch_idx to int64 before multiplication to prevent compilation error.
    rand_offset = batch_idx.to(tl.int64) * 2654435761 # A large prime for better hash
    random_uniform = tl.rand(rand_seed, rand_offset)

    # FIX: Replace incorrect serial for-loop with a fully vectorized sampling implementation.
    # 1. Compute the Cumulative Distribution Function (CDF).
    sample_cdf = tl.cumsum(final_probs, axis=0)

    # 2. Find the first index where the random number is less than the CDF.
    # This creates a mask like [False, False, True, True, ...].
    sampling_mask = (random_uniform < sample_cdf) & (final_probs > 0.0)

    # 3. Find the minimum index where this mask is True.
    # Where the mask is False, replace the index with a large value.
    masked_arange = tl.where(sampling_mask, k_arange, VOCAB_SIZE_P2)
    # The minimum value of this tensor is the relative index we want.
    sampled_arange_idx = tl.min(masked_arange, axis=0)

    # 4. Use the found index to look up the actual token ID from `sorted_indices`.
    # This is a gather operation where the index is a scalar.
    lookup_mask = (k_arange == sampled_arange_idx)
    sampling_sample_tensor = tl.where(lookup_mask, sorted_indices, 0)
    sampling_sample = tl.sum(sampling_sample_tensor, axis=0)
    
    # Fallback: if all filtered probabilities were zero, the min index will be VOCAB_SIZE_P2.
    # In this case, we default to the greedy sample.
    all_probs_zero = (sampled_arange_idx == VOCAB_SIZE_P2)
    sampling_sample = tl.where(all_probs_zero, greedy_sample, sampling_sample)

    # --- 6. Final Selection and Store ---
    # If p <= 0.0, use the greedy sample. Otherwise, use the sampled result.
    final_sample = tl.where(p <= 0.0, greedy_sample, sampling_sample)
    
    # Store the final sampled token index, casting to the required int64 type.
    tl.store(samples_ptr + batch_idx, final_sample.to(tl.int64))


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

    This function handles device management, dtype conversions, and kernel launch.
    It can be called with positional or keyword arguments.

    Args:
        probs (torch.Tensor): Probability distributions of shape [batch_size, vocab_size].
        top_k (torch.Tensor): Top-k values for each sequence of shape [batch_size].
        top_p (torch.Tensor): Top-p values for each sequence of shape [batch_size].

    Returns:
        torch.Tensor: Sampled token indices of shape [batch_size].
    """
    # --- 0. Argument Parsing ---
    if args:
        if len(args) != 3:
            raise ValueError(f"Expected 3 positional arguments (probs, top_k, top_p), but got {len(args)}")
        probs, top_k, top_p = args
    else:
        try:
            probs = kwargs["probs"]
            top_k = kwargs["top_k"]
            top_p = kwargs["top_p"]
        except KeyError as e:
            raise KeyError(f"Missing required keyword argument: {e}") from e

    # --- 1. Validation and Device Management ---
    if not isinstance(probs, torch.Tensor) or probs.dim() != 2:
        raise ValueError(f"Input 'probs' must be a 2D torch.Tensor, but got {type(probs)}")
    if not isinstance(top_k, torch.Tensor) or top_k.dim() != 1:
        raise ValueError(f"Input 'top_k' must be a 1D torch.Tensor, but got {type(top_k)}")
    if not isinstance(top_p, torch.Tensor) or top_p.dim() != 1:
        raise ValueError(f"Input 'top_p' must be a 1D torch.Tensor, but got {type(top_p)}")

    batch_size, vocab_size = probs.shape
    if top_k.shape[0] != batch_size or top_p.shape[0] != batch_size:
        raise ValueError("Batch dimensions of all inputs must match.")
    
    if vocab_size != 129280:
        raise ValueError(f"vocab_size must be 129280, but got {vocab_size}")

    original_device = probs.device
    if torch.cuda.is_available():
        device = torch.device("cuda")
    else:
        if original_device.type == 'cpu':
            raise RuntimeError("This implementation requires a CUDA-enabled GPU, but input tensors are on CPU.")
        device = original_device

    # Ensure all tensors are on the same CUDA device and have the correct dtype
    probs = probs.to(device=device, dtype=torch.float32)
    top_k = top_k.to(device=device, dtype=torch.int32)
    top_p = top_p.to(device=device, dtype=torch.float32)
    
    # --- 2. Kernel Configuration ---
    samples = torch.empty(batch_size, dtype=torch.int64, device=device)
    
    vocab_size_p2 = triton.next_power_of_2(vocab_size)

    grid = (batch_size,)
    
    rand_seed = torch.randint(0, 2**32 - 1, (1,), device='cpu').item()

    # --- 3. Kernel Launch ---
    top_k_top_p_sampling_from_probs_v129280_kernel[grid](
        probs_ptr=probs,
        top_k_ptr=top_k,
        top_p_ptr=top_p,
        samples_ptr=samples,
        rand_seed=rand_seed,
        VOCAB_SIZE=vocab_size,
        VOCAB_SIZE_P2=vocab_size_p2,
    )

    # --- 4. Return to Original Device ---
    if samples.device != original_device:
        samples = samples.to(original_device)
        
    return samples
scrolls · 247 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON