Skip to content
KernelIndex
Search⌘K

gemini-2.5-pro / triton2a8f55

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

Use it

Vendorable · source mirrored · Apache-2.0View source →

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

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

Kernel source

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

@triton.jit
def _top_k_sampling_from_probs_kernel(
    probs_ptr,       # Pointer to [batch_size, vocab_size] float32 tensor
    top_k_ptr,       # Pointer to [batch_size] int32 tensor
    samples_ptr,     # Pointer to [batch_size] int64 output tensor
    seed,            # Scalar uint64 seed for random number generation
    batch_size,      # Number of sequences in the batch
    VOCAB_SIZE: tl.constexpr,
    BLOCK_SIZE_V: tl.constexpr,
):
    """
    Triton kernel for top-k sampling from probability distributions.
    This kernel avoids allocating large intermediate tensors by using a multi-pass
    approach to find the top-k threshold and perform sampling in a memory-efficient manner.

    Strategy:
    1. For each sequence, determine if top-k filtering is necessary (i.e., 0 < k < vocab_size).
    2. If filtering is on:
       a. Find the k-th largest probability value (the threshold) using binary search over
          the probability values. This involves multiple passes but is memory-efficient.
       b. Calculate the sum of probabilities strictly greater than the threshold (sum_gt) and
          the count of such probabilities (count_gt).
       c. The total sum for the new distribution is sum_gt plus the sum of (k - count_gt)
          elements that are equal to the threshold.
    3. If filtering is off (k is invalid or covers the full vocab):
       a. The total sum is simply the sum of all probabilities.
    4. A random number is generated and scaled by the total_sum to determine a target value.
    5. A final vectorized scan over the vocabulary applies the filtering logic on-the-fly
       and uses a cumulative sum approach to find the token index corresponding to the target value.
    """
    # Each program instance processes one sequence from the batch.
    pid = tl.program_id(axis=0)

    # Pointers for the current sequence
    row_probs_ptr = probs_ptr + pid * VOCAB_SIZE
    row_top_k_ptr = top_k_ptr + pid
    row_samples_ptr = samples_ptr + pid

    k = tl.load(row_top_k_ptr)

    # =================================================================
    # Step 1: Determine threshold and sum for sampling
    # =================================================================
    
    threshold = -1.0
    total_sum = 0.0
    sum_gt = 0.0
    count_gt = 0
    
    do_filter = (k > 0) & (k < VOCAB_SIZE)

    if do_filter:
        # --- Pass 1: Find max probability to bound the binary search ---
        max_prob = 0.0
        for offset in range(0, VOCAB_SIZE, BLOCK_SIZE_V):
            v_offsets = offset + tl.arange(0, BLOCK_SIZE_V)
            mask = v_offsets < VOCAB_SIZE
            p = tl.load(row_probs_ptr + v_offsets, mask=mask, other=0.0)
            chunk_max = tl.max(p, axis=0)
            max_prob = tl.maximum(max_prob, chunk_max)
        
        # --- Pass 2: Binary search for the k-th probability value (threshold) ---
        low = 0.0
        high = max_prob
        # 16 iterations are sufficient for float32 precision
        for _ in range(16):
            mid = 0.5 * (low + high)
            count = 0
            for offset in range(0, VOCAB_SIZE, BLOCK_SIZE_V):
                v_offsets = offset + tl.arange(0, BLOCK_SIZE_V)
                mask = v_offsets < VOCAB_SIZE
                p = tl.load(row_probs_ptr + v_offsets, mask=mask, other=0.0)
                # FIX: Use >= to correctly find the k-th value as the threshold,
                # which is crucial for handling cases with duplicate probability values.
                count += tl.sum((p >= mid).to(tl.int32), axis=0)
            
            if count >= k:
                low = mid
            else:
                high = mid
        threshold = low

        # --- Pass 3: Calculate sum of probs > threshold and count of probs > threshold ---
        for offset in range(0, VOCAB_SIZE, BLOCK_SIZE_V):
            v_offsets = offset + tl.arange(0, BLOCK_SIZE_V)
            mask = v_offsets < VOCAB_SIZE
            p = tl.load(row_probs_ptr + v_offsets, mask=mask, other=0.0)
            
            is_gt = p > threshold
            sum_gt += tl.sum(tl.where(is_gt, p, 0.0), axis=0)
            count_gt += tl.sum(is_gt.to(tl.int32), axis=0)

        k_rem = k - count_gt
        k_rem = tl.maximum(0, k_rem)
        
        sum_eq = k_rem.to(tl.float32) * threshold
        total_sum = sum_gt + sum_eq
    
    else: # k is invalid or full vocab, sample from the original distribution
        for offset in range(0, VOCAB_SIZE, BLOCK_SIZE_V):
            v_offsets = offset + tl.arange(0, BLOCK_SIZE_V)
            mask = v_offsets < VOCAB_SIZE
            p = tl.load(row_probs_ptr + v_offsets, mask=mask, other=0.0)
            total_sum += tl.sum(p, axis=0)

    # Handle cases where the distribution sum is zero to sample uniformly.
    if total_sum <= 1e-9:
        rand_offset = pid.to(tl.uint64)
        rand_uint32 = tl.rand(seed, rand_offset)
        # FIX: Use float multiplication to avoid modulo bias for uniform sampling
        rand_float = (rand_uint32 / 4294967296.0).to(tl.float32)
        rand_idx = (rand_float * VOCAB_SIZE).to(tl.int32)
        # Clamp to ensure index is within bounds
        rand_idx = tl.minimum(rand_idx, VOCAB_SIZE - 1)
        tl.store(row_samples_ptr, rand_idx.to(tl.int64))
        return

    # =================================================================
    # Step 2: Multinomial Sampling Scan (Vectorized)
    # =================================================================

    rand_offset = pid.to(tl.uint64) + VOCAB_SIZE # Use a different offset for this random number
    rand_uint32 = tl.rand(seed, rand_offset)
    # Scale uint32 random int to a float32 in [0, 1)
    rand_float = (rand_uint32 / 4294967296.0).to(tl.float32)
    sample_val = rand_float * total_sum
    
    # Initialize final_idx with a Python int. Triton infers its type as tl.int32.
    final_idx = VOCAB_SIZE
    
    is_in_gt_bucket = sample_val < sum_gt
    
    cumsum = 0.0
    eq_count = 0
    
    target_eq_idx = 0
    if do_filter and not is_in_gt_bucket:
        target_eq_rem = sample_val - sum_gt
        safe_threshold = tl.where(threshold > 0.0, threshold, 1.0)
        target_eq_idx = tl.floor(target_eq_rem / safe_threshold).to(tl.int32)

    for offset in range(0, VOCAB_SIZE, BLOCK_SIZE_V):
        v_offsets = offset + tl.arange(0, BLOCK_SIZE_V)
        load_mask = v_offsets < VOCAB_SIZE
        p = tl.load(row_probs_ptr + v_offsets, mask=load_mask, other=0.0)

        match_indices = tl.full((BLOCK_SIZE_V,), VOCAB_SIZE, dtype=tl.int32)

        if not do_filter:
            block_cumsum = cumsum + tl.cumsum(p, axis=0)
            is_match = (sample_val < block_cumsum) & (sample_val >= (block_cumsum - p))
            match_indices = tl.where(is_match & load_mask, v_offsets, VOCAB_SIZE)
            cumsum += tl.sum(p, axis=0)
        else:
            if is_in_gt_bucket:
                filtered_p = tl.where(p > threshold, p, 0.0)
                block_cumsum = cumsum + tl.cumsum(filtered_p, axis=0)
                is_match = (sample_val < block_cumsum) & (sample_val >= (block_cumsum - filtered_p))
                match_indices = tl.where(is_match & load_mask, v_offsets, VOCAB_SIZE)
                cumsum += tl.sum(filtered_p, axis=0)
            else:
                is_eq = p == threshold
                block_eq_cumsum = eq_count + tl.cumsum(is_eq.to(tl.int32), axis=0)
                is_match = (block_eq_cumsum == (target_eq_idx + 1)) & is_eq
                match_indices = tl.where(is_match & load_mask, v_offsets, VOCAB_SIZE)
                eq_count += tl.sum(is_eq.to(tl.int32), axis=0)
        
        block_min_idx = tl.min(match_indices, axis=0)
        # Keep all operations in int32 to maintain type consistency for the
        # loop-carried variable 'final_idx'.
        final_idx = tl.minimum(final_idx, block_min_idx)

    # If no index was found (e.g., due to floating point rounding), default to the last valid index.
    final_idx = tl.where(final_idx >= VOCAB_SIZE, VOCAB_SIZE - 1, final_idx)
    # Cast to int64 at the very end to match the output tensor's dtype.
    tl.store(row_samples_ptr, final_idx.to(tl.int64))


@torch.no_grad()
def _reference_run(probs, top_k):
    """
    Reference PyTorch implementation for functionality verification and CPU fallback.
    This version is careful to not modify the input tensor in-place.
    """
    batch_size, vocab_size = probs.shape
    device = probs.device
    assert vocab_size == 128256

    probs_float = probs.to(torch.float32)
    samples = torch.empty(batch_size, dtype=torch.int64, device=device)

    for i in range(batch_size):
        row = probs_float[i]
        k = int(top_k[i].item())
        
        sampling_dist = row

        if 0 < k < vocab_size:
            idx_sorted = torch.argsort(row, descending=True)
            keep_idx = idx_sorted[:k]

            filtered = torch.zeros_like(row)
            filtered[keep_idx] = row[keep_idx]
            
            # If the sum of top-k probabilities is positive, sample from them.
            if filtered.sum() > 1e-9:
                sampling_dist = filtered
            else:
                sampling_dist = torch.ones_like(row)
        
        # Ensure we don't pass a zero-sum tensor to multinomial, which would error on some torch versions.
        # It's specified to sample uniformly in that case.
        if sampling_dist.sum() <= 1e-9:
            sampling_dist = torch.ones_like(row)

        samples[i] = torch.multinomial(sampling_dist, 1, replacement=True).squeeze(0)

    return samples


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

    Handles device management, argument parsing, grid computation, and error checking.
    It preserves the device of the input tensors for the output.

    Args:
        probs (torch.Tensor): A [batch_size, vocab_size] tensor of float32 probabilities.
        top_k (torch.Tensor): A [batch_size] tensor of int32 values for k.
        
    Returns:
        torch.Tensor: A [batch_size] tensor of int64 sampled token indices.
    """
    # 1. Argument Parsing
    if args:
        if len(args) > 2:
            raise ValueError(f"Expected 2 positional arguments, but got {len(args)}")
        probs, top_k = args
    else:
        probs = kwargs.get("probs")
        top_k = kwargs.get("top_k")

    if probs is None or top_k is None:
        raise ValueError("Missing required arguments 'probs' and 'top_k'")

    # 2. Input Validation
    if not isinstance(probs, torch.Tensor) or not isinstance(top_k, torch.Tensor):
        raise TypeError("Inputs 'probs' and 'top_k' must be torch.Tensors.")
        
    if probs.ndim != 2:
        raise ValueError(f"Input 'probs' must be a 2D tensor, but got shape {probs.shape}")
    if top_k.ndim != 1:
        raise ValueError(f"Input 'top_k' must be a 1D tensor, but got shape {top_k.shape}")
    
    batch_size, vocab_size = probs.shape
    if top_k.shape[0] != batch_size:
        raise ValueError(f"Dimension mismatch: probs.shape[0] ({batch_size}) != top_k.shape[0] ({top_k.shape[0]})")
    
    VOCAB_SIZE = 128256
    if vocab_size != VOCAB_SIZE:
        raise ValueError(f"vocab_size must be {VOCAB_SIZE}, but got {vocab_size}")

    # 3. Device Management
    original_device = probs.device
    
    if not torch.cuda.is_available():
        if original_device.type != 'cpu':
             raise RuntimeError("CUDA is not available, but input tensors are on a CUDA device.")
        # Fallback to reference implementation on CPU if CUDA is not available
        return _reference_run(probs, top_k)
    
    device = torch.device('cuda')
    # Move inputs to the default CUDA device if they aren't already there
    probs = probs.to(device)
    top_k = top_k.to(device)

    # Ensure contiguous inputs and correct dtypes for the kernel
    probs = probs.contiguous().to(torch.float32)
    top_k = top_k.contiguous().to(torch.int32)
    
    # 4. Grid and Kernel Execution
    samples = torch.empty(batch_size, dtype=torch.int64, device=device)
    
    # Use a large block size for the vocabulary dimension to maximize memory-level parallelism,
    # which is crucial for this memory-bound kernel, especially on modern GPUs like B200.
    BLOCK_SIZE_V = 2048

    grid = (batch_size,)

    # Generate a random seed for the kernel for reproducibility
    seed = torch.randint(0, 2**63 - 1, (1,)).item()
    
    _top_k_sampling_from_probs_kernel[grid](
        probs_ptr=probs,
        top_k_ptr=top_k,
        samples_ptr=samples,
        seed=seed,
        batch_size=batch_size,
        VOCAB_SIZE=VOCAB_SIZE,
        BLOCK_SIZE_V=BLOCK_SIZE_V,
    )
    
    # 5. Output Device Management
    # Move the result back to the original device if necessary
    if samples.device != original_device:
        samples = samples.to(original_device)
        
    return samples
scrolls · 314 lines total

Source code from the importing source · Apache-2.0

No published measurement for this revision

JSON