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