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 samplesscrolls · 247 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON