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