gemini-2.5-pro / tritonf8ce0a
gemini-2.5-pro_triton_f8ce0a · gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 264 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-f8ce0a?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:e3706a7eb0b412bf452df3285c737486b0efce400e39f9cb6002a57765d900f6
license declaredApache-2.0
license concludedApache-2.0
authorsgemini-2.5-pro
imported2026-08-20
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
@triton.autotune(num-warps = 4
triton.Config({'BLOCK_V': 1024, 'MAX_K_BUFFER': 2048}, num_warps=4),Kernel source
main.py264 lines
import torch
import triton
import triton.language as tl
import math
# `triton.jit`'ed functions can be auto-tuned by using the `triton.autotune` decorator, which consumes:
# - A list of `triton.Config` objects that define different configurations of
# meta-parameters (e.g., `BLOCK_SIZE_M`) and compiler options (e.g., `num_warps`)
# - A `key` argument containing argument names for the kernel parameters
#
# JITed functions can be decorated with `triton.autotune` to optimize for a given input shape.
# This is especially important for kernels that handle tensors with variable shapes.
# The `key` argument is used to lookup the best configuration for a given set of input shapes.
# For this kernel, we tune for batch_size.
@triton.autotune(
configs=[
triton.Config({'BLOCK_V': 1024, 'MAX_K_BUFFER': 2048}, num_warps=4),
triton.Config({'BLOCK_V': 2048, 'MAX_K_BUFFER': 2048}, num_warps=8),
triton.Config({'BLOCK_V': 1024, 'MAX_K_BUFFER': 4096}, num_warps=4),
triton.Config({'BLOCK_V': 2048, 'MAX_K_BUFFER': 4096}, num_warps=8),
triton.Config({'BLOCK_V': 512, 'MAX_K_BUFFER': 2048}, num_warps=2),
],
key=['batch_size'],
)
@triton.jit
def top_k_top_p_sampling_from_probs_v151936_kernel(
# Pointers to tensors
probs_ptr,
top_k_ptr,
top_p_ptr,
samples_ptr,
# Random sampling state
seed,
offsets_ptr,
# Tensor dimensions
batch_size,
# Strides
stride_probs_b,
# Meta-parameters
VOCAB_SIZE: tl.constexpr,
BLOCK_V: tl.constexpr,
MAX_K_BUFFER: tl.constexpr,
):
"""
Triton kernel for top-k/top-p sampling.
Each program instance processes one sequence from the batch.
"""
# -----------------------------------------------------------
# Program setup
# -----------------------------------------------------------
pid = tl.program_id(0)
# Load per-sequence sampling parameters
k = tl.load(top_k_ptr + pid)
p = tl.load(top_p_ptr + pid)
# Pointers to the current sequence's data
probs_row_ptr = probs_ptr + pid * stride_probs_b
# -----------------------------------------------------------
# Greedy decoding path (p <= 0.0) -> argmax
# -----------------------------------------------------------
if p <= 0.0:
max_prob = -1.0
max_idx = -1
v_offsets = tl.arange(0, BLOCK_V)
for v_start in range(0, VOCAB_SIZE, BLOCK_V):
v_range = v_start + v_offsets
v_mask = v_range < VOCAB_SIZE
row_probs = tl.load(probs_row_ptr + v_range, mask=v_mask, other=-1.0)
block_max_prob = tl.max(row_probs)
# If the max in this block is greater than the global max, update global max
# and find the first index of this new max in the current block.
if block_max_prob > max_prob:
max_prob = block_max_prob
is_max = (row_probs == max_prob) & v_mask
max_indices_in_block = tl.where(is_max, v_range, VOCAB_SIZE + 1)
max_idx = tl.min(max_indices_in_block)
# If the block max is equal to the global max, we only update the index
# if the new index is smaller (torch.argmax behavior).
elif block_max_prob == max_prob:
is_max = (row_probs == max_prob) & v_mask
max_indices_in_block = tl.where(is_max, v_range, VOCAB_SIZE + 1)
block_min_idx = tl.min(max_indices_in_block)
if block_min_idx < max_idx:
max_idx = block_min_idx
tl.store(samples_ptr + pid, max_idx)
return
# -----------------------------------------------------------
# Top-K and Top-P Sampling Path
# -----------------------------------------------------------
# --- Stage 1: Find top candidates using a streaming approach ---
# `effective_k` is the number of candidates to consider after sorting.
# If k is invalid (<=0) or too large, we default to the buffer size for candidate search,
# but k will be respected during filtering.
effective_k = k
if k <= 0 or k > MAX_K_BUFFER:
effective_k = MAX_K_BUFFER
# Initialize SRAM buffers with the first block of candidates.
v_offsets_init = tl.arange(0, MAX_K_BUFFER)
v_mask_init = v_offsets_init < VOCAB_SIZE
sram_probs = tl.load(probs_row_ptr + v_offsets_init, mask=v_mask_init, other=-1.0)
sram_indices = v_offsets_init.to(tl.int32)
min_prob_in_sram = tl.min(sram_probs)
# Iterate over the rest of the vocabulary to find better candidates.
v_offsets = tl.arange(0, BLOCK_V)
for v_start in range(MAX_K_BUFFER, VOCAB_SIZE, BLOCK_V):
v_range = v_start + v_offsets
v_mask = v_range < VOCAB_SIZE
block_probs = tl.load(probs_row_ptr + v_range, mask=v_mask, other=-1.0)
# Optimization: only process block if it contains a potential candidate
if tl.max(block_probs) > min_prob_in_sram:
# This loop is unrolled by the compiler. It serially updates the candidate set.
for i in range(BLOCK_V):
prob = tl.load(probs_row_ptr + v_start + i, mask=(v_start + i < VOCAB_SIZE), other=-1.0)
if prob > min_prob_in_sram:
# Find the location of the minimum element and replace it.
min_mask = sram_probs == min_prob_in_sram
# To break ties, take the one with the smallest index in the sram buffer.
min_indices = tl.where(min_mask, tl.arange(0, MAX_K_BUFFER), MAX_K_BUFFER + 1)
first_min_idx = tl.min(min_indices)
# Replace the minimum element with the new, larger candidate.
sram_probs = tl.where(tl.arange(0, MAX_K_BUFFER) == first_min_idx, prob, sram_probs)
sram_indices = tl.where(tl.arange(0, MAX_K_BUFFER) == first_min_idx, v_start + i, sram_indices)
# Update the minimum for the next iteration of this inner loop.
min_prob_in_sram = tl.min(sram_probs)
# Sort the final candidates before top-p filtering.
# We use a robust packing method to sort key-value pairs.
packed = (sram_probs * 2147483647.0).to(tl.int32).to(tl.int64) << 32 | sram_indices.to(tl.int64)
sorted_packed = tl.sort(packed, descending=True)
# Corrected unpacking: Do not mask the sign bit.
sram_probs = (sorted_packed >> 32).to(tl.int32).to(tl.float32) / 2147483647.0
sram_indices = (sorted_packed & 0xFFFFFFFF).to(tl.int32)
# --- Stage 2: Apply Top-K then Top-P filtering on the candidates ---
k_arange = tl.arange(0, MAX_K_BUFFER)
k_mask = k_arange < effective_k
masked_sram_probs = tl.where(k_mask, sram_probs, 0.0)
total_prob_k = tl.sum(masked_sram_probs, axis=0)
norm_sram_probs = masked_sram_probs / (total_prob_k + 1e-9)
cumsum_probs = tl.cumsum(norm_sram_probs, axis=0)
# Condition for discarding token `i` is when cumulative prob of tokens `0..i-1` >= `p`.
p_mask = (cumsum_probs - norm_sram_probs) >= p
p_cutoff_idx_raw = tl.where(p_mask, k_arange, MAX_K_BUFFER)
num_final_candidates = tl.min(p_cutoff_idx_raw)
# Ensure at least one token is considered, matching reference behavior.
if num_final_candidates == 0:
num_final_candidates = 1
# The number of final candidates cannot exceed the top-k limit.
if num_final_candidates > effective_k:
num_final_candidates = effective_k
final_mask = k_arange < num_final_candidates
final_probs = tl.where(final_mask, norm_sram_probs, 0.0)
total_prob_p = tl.sum(final_probs, axis=0)
# --- Stage 3: Sample from the final candidates ---
rand_offset = tl.load(offsets_ptr + pid)
tl.store(offsets_ptr + pid, rand_offset + 1)
random_uniform = tl.rand(seed, rand_offset)
random_scaled = random_uniform * (total_prob_p + 1e-9)
final_cumsum = tl.cumsum(final_probs, axis=0)
sampled_mask = random_scaled < final_cumsum
sampled_idx_in_sram_raw = tl.where(sampled_mask, k_arange, MAX_K_BUFFER)
sampled_idx_in_sram = tl.min(sampled_idx_in_sram_raw)
# Gather the final token index from the sram buffer.
selection_mask = k_arange == sampled_idx_in_sram
final_sample_idx = tl.sum(tl.where(selection_mask, sram_indices, 0))
tl.store(samples_ptr + pid, final_sample_idx.to(tl.int64))
def run(probs: torch.Tensor, top_k: torch.Tensor, top_p: torch.Tensor, **kwargs):
"""
Wrapper function for the top_k_top_p_sampling Triton kernel.
Args:
probs (torch.Tensor): Probability distributions [batch_size, vocab_size], DType.FLOAT32.
top_k (torch.Tensor): Number of top tokens to consider [batch_size], DType.INT32.
top_p (torch.Tensor): Cumulative probability threshold [batch_size], DType.FLOAT32.
Returns:
torch.Tensor: Sampled token indices [batch_size], DType.INT64.
"""
# -----------------------------------------------------------
# Device and DType management
# -----------------------------------------------------------
original_device = probs.device
if not torch.cuda.is_available():
if any(t.is_cuda for t in [probs, top_k, top_p]):
raise RuntimeError("CUDA is required for this Triton kernel, but is not available.")
# This path is for CPU-only environments, which Triton doesn't support.
raise RuntimeError("CUDA is required for this Triton kernel.")
compute_device = torch.device('cuda')
# Move all tensors to the compute device
probs = probs.to(compute_device, non_blocking=True)
top_k = top_k.to(compute_device, non_blocking=True)
top_p = top_p.to(compute_device, non_blocking=True)
# Ensure correct dtypes
probs = probs.to(torch.float32)
top_k = top_k.to(torch.int32)
top_p = top_p.to(torch.float32)
# -----------------------------------------------------------
# Kernel launch setup
# -----------------------------------------------------------
batch_size, vocab_size = probs.shape
if vocab_size != 151936:
raise ValueError(f"This kernel is specialized for vocab_size=151936, but got {vocab_size}")
samples = torch.empty(batch_size, dtype=torch.int64, device=compute_device)
# Seed and offsets for random number generation
seed = 1234
offsets = torch.randint(0, vocab_size * 2, (batch_size,), dtype=torch.int32, device=compute_device)
grid = lambda meta: (batch_size,)
# -----------------------------------------------------------
# Kernel invocation
# -----------------------------------------------------------
top_k_top_p_sampling_from_probs_v151936_kernel[grid](
probs,
top_k,
top_p,
samples,
seed,
offsets,
batch_size,
probs.stride(0),
VOCAB_SIZE=vocab_size
)
# -----------------------------------------------------------
# Finalization
# -----------------------------------------------------------
# Move the result back to the original device
return samples.to(original_device)scrolls · 264 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON