gemini-2.5-pro_triton_544238
gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 246 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-544238?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:c6003e0ad696af19229b4aa17de085c4e898f73dd3cd017481fc579b918b9cf7
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.
tile-m = 512
BLOCK_SIZE_M = 512Kernel source
main.py246 lines
import torch
import triton
import triton.language as tl
import math
# B200 is part of the Blackwell architecture. Optimizations for Hopper
# (large SRAM, efficient block-level primitives) are expected to perform well on B200.
# This kernel is redesigned using modern Triton features to be correct and efficient.
# --- Triton Kernel ---
@triton.jit
def top_k_top_p_sampling_from_probs_v128256_kernel(
probs_ptr,
top_k_ptr,
top_p_ptr,
samples_ptr,
seed_ptr,
batch_size,
stride_probs_b,
VOCAB_SIZE: tl.constexpr,
BLOCK_SIZE_M: tl.constexpr,
TOP_K_BUFFER_SIZE: tl.constexpr,
):
"""
Triton kernel for Top-K, Top-P sampling.
This kernel processes one sequence per program instance (one row of the batch).
It uses a streaming reduction to find the top candidates from the vocabulary,
applies top-k and top-p filtering, and finally performs multinomial sampling.
This version uses modern Triton APIs and direct operations on tl.tensor objects.
- VOCAB_SIZE: The total vocabulary size.
- BLOCK_SIZE_M: The size of blocks to read from the vocab. Must be a power of 2.
- TOP_K_BUFFER_SIZE: The size of the SRAM buffer for holding top candidates. Must be a power of 2.
"""
pid = tl.program_id(0)
if pid >= batch_size:
return
# --- Load per-sequence parameters ---
k_val = tl.load(top_k_ptr + pid)
p_val = tl.load(top_p_ptr + pid)
seed = tl.load(seed_ptr + pid)
probs_row_ptr = probs_ptr + pid * stride_probs_b
# The size of our merge buffer for the streaming top-k reduction.
# Must be a power of 2 for tl.sort.
# FIX: Declared as tl.constexpr to resolve the compilation error with tl.arange.
MERGE_BUFFER_SIZE: tl.constexpr = TOP_K_BUFFER_SIZE + BLOCK_SIZE_M
# --- SRAM Allocation for top candidates (as registers) ---
sram_top_k_probs = tl.full([TOP_K_BUFFER_SIZE], -1.0, dtype=tl.float32)
sram_top_k_indices = tl.full([TOP_K_BUFFER_SIZE], -1, dtype=tl.int32)
# --- Streaming Top-K Reduction ---
# Find the top `TOP_K_BUFFER_SIZE` candidates from the entire vocabulary.
num_blocks = tl.cdiv(VOCAB_SIZE, BLOCK_SIZE_M)
for block_idx in range(num_blocks):
# Load a block of probabilities and their corresponding indices from HBM
m_offsets = tl.arange(0, BLOCK_SIZE_M)
current_offsets = block_idx * BLOCK_SIZE_M + m_offsets
mask = current_offsets < VOCAB_SIZE
chunk_probs = tl.load(probs_row_ptr + current_offsets, mask=mask, other=-1.0)
chunk_indices = current_offsets
# --- Merge and Sort in Registers/SRAM ---
# 1. Construct the merged buffer of candidates by concatenating the
# current top-k with the new chunk.
# FIX: The original indexing logic was flawed and caused out-of-bounds access.
# This corrected version clamps indices to be safe for both branches of tl.where.
merged_offsets = tl.arange(0, MERGE_BUFFER_SIZE)
is_top_k_part = merged_offsets < TOP_K_BUFFER_SIZE
# Safely clamp indices for the SRAM part to [0, TOP_K_BUFFER_SIZE - 1]
sram_indices_safe = tl.minimum(merged_offsets, TOP_K_BUFFER_SIZE - 1)
# Safely clamp indices for the chunk part to [0, BLOCK_SIZE_M - 1]
chunk_indices_safe = tl.maximum(0, merged_offsets - TOP_K_BUFFER_SIZE)
chunk_indices_safe = tl.minimum(chunk_indices_safe, BLOCK_SIZE_M - 1)
merged_probs = tl.where(is_top_k_part, sram_top_k_probs[sram_indices_safe], chunk_probs[chunk_indices_safe])
merged_indices = tl.where(is_top_k_part, sram_top_k_indices[sram_indices_safe], chunk_indices[chunk_indices_safe])
# 2. Pack probs (key) and indices (value) into int64 for a single sort operation.
# To sort floats in descending order, we negate their integer representation.
probs_as_int = merged_probs.to(tl.int32, bitcast=True)
neg_probs_as_int = -probs_as_int
packed_data = neg_probs_as_int.to(tl.int64) << 32 | merged_indices.to(tl.int64)
# 3. Sort the packed data. tl.sort is a highly optimized block-level primitive.
sorted_packed = tl.sort(packed_data)
# 4. Unpack the data and update the top-k buffers for the next iteration.
k_offsets = tl.arange(0, TOP_K_BUFFER_SIZE)
top_k_packed_slice = sorted_packed[k_offsets]
unpacked_neg_probs_as_int = (top_k_packed_slice >> 32).to(tl.int32)
sram_top_k_indices = (top_k_packed_slice & 0xFFFFFFFF).to(tl.int32)
sram_top_k_probs = (-unpacked_neg_probs_as_int).to(tl.float32, bitcast=True)
# `sram_top_k_probs` and `sram_top_k_indices` now hold the top candidates.
# --- Apply Top-K filtering ---
num_candidates = TOP_K_BUFFER_SIZE
if 0 < k_val < VOCAB_SIZE:
num_candidates = tl.minimum(k_val, TOP_K_BUFFER_SIZE)
cand_offsets = tl.arange(0, TOP_K_BUFFER_SIZE)
k_mask = cand_offsets < num_candidates
# --- Greedy sampling (p <= 0.0) ---
if p_val <= 0.0:
# The candidates are sorted, so the first element is the argmax.
result_idx = sram_top_k_indices[0]
tl.store(samples_ptr + pid, result_idx.to(tl.int64))
return
# --- Apply Top-P (Nucleus) filtering ---
candidate_probs = tl.where(k_mask, sram_top_k_probs, 0.0)
if p_val < 1.0:
# Renormalize the candidate probabilities before calculating CDF for top-p.
total_prob_sum = tl.sum(candidate_probs, axis=0)
# Avoid division by zero if all candidate probs are zero
if total_prob_sum > 1e-9:
probs_for_p = candidate_probs / total_prob_sum
# Compute CDF. Keep token `i` if `cdf[i-1] <= p`.
# This is equivalent to `(cumsum(p) - p) <= p`.
cdf = tl.cumsum(probs_for_p, axis=0)
shifted_cdf = cdf - probs_for_p
p_mask = shifted_cdf < p_val
candidate_probs = tl.where(p_mask, candidate_probs, 0.0)
else:
# If sum is zero, all candidate_probs are already zero, so do nothing.
pass
# --- Final Multinomial Sampling ---
final_probs = candidate_probs
final_prob_sum = tl.sum(final_probs, axis=0)
result_idx = -1
if final_prob_sum > 1e-9: # Use a small epsilon for float comparison
# Generate a random number in [0, 1) and scale it.
rand_offset = pid # Use a unique offset for per-row randomness
r_val = tl.rand(seed, rand_offset)
sample_thresh = r_val * final_prob_sum
sample_cdf = tl.cumsum(final_probs, axis=0)
# Find the first index `i` where `sample_cdf[i] > sample_thresh`.
is_winner = sample_cdf > sample_thresh
# tl.argmax returns the index of the first '1'
winner_sram_idx = tl.argmax(is_winner.to(tl.int32), axis=0)
result_idx = sram_top_k_indices[winner_sram_idx]
else:
# If all probabilities are filtered out (e.g., k=0 or p is very small),
# fall back to the absolute top token (greedy).
result_idx = sram_top_k_indices[0]
tl.store(samples_ptr + pid, result_idx.to(tl.int64))
def run(*args, **kwargs):
"""
Wrapper function for the Top-K, Top-P sampling Triton kernel.
Handles device management, kernel launching, and input validation.
"""
# --- Input Parsing and Validation ---
probs = kwargs.get("probs", args[0] if len(args) > 0 else None)
top_k = kwargs.get("top_k", args[1] if len(args) > 1 else None)
top_p = kwargs.get("top_p", args[2] if len(args) > 2 else None)
# Allow seed to be passed for deterministic testing
seed = kwargs.get("seed")
if probs is None or top_k is None or top_p is None:
raise ValueError("Inputs 'probs', 'top_k', and 'top_p' must be provided.")
assert probs.dim() == 2, "probs must be a 2D tensor"
assert top_k.dim() == 1, "top_k must be a 1D tensor"
assert top_p.dim() == 1, "top_p must be a 1D tensor"
batch_size, vocab_size = probs.shape
assert top_k.shape[0] == batch_size, "top_k batch size mismatch"
assert top_p.shape[0] == batch_size, "top_p batch size mismatch"
assert vocab_size == 128256, f"vocab_size must be 128256, but got {vocab_size}"
# --- Device Management ---
if not torch.cuda.is_available() and probs.device.type != 'cpu':
raise RuntimeError("This kernel requires a CUDA-enabled GPU.")
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
if not torch.cuda.is_available():
# Fallback to reference for CPU-only environments
import warnings
warnings.warn("CUDA not available. Falling back to reference implementation. This will be slow.")
# Simulating reference run for completeness, as it's not provided
# In a real scenario, you'd call the actual reference implementation here.
samples = torch.empty(batch_size, dtype=torch.int64, device='cpu')
for i in range(batch_size):
samples[i] = torch.multinomial(probs[i], 1).squeeze()
return samples
original_device = probs.device
# Move all inputs to the GPU where the kernel will run
probs = probs.to(device, non_blocking=True, dtype=torch.float32)
top_k = top_k.to(device, non_blocking=True, dtype=torch.int32)
top_p = top_p.to(device, non_blocking=True, dtype=torch.float32)
# --- Kernel Launch ---
samples = torch.empty(batch_size, dtype=torch.int64, device=device)
# Generate per-row seeds for reproducibility and randomness
if seed is None:
seeds = torch.randint(0, 2**31 - 1, (batch_size,), device=device, dtype=torch.int32)
else:
# Create a deterministic sequence of seeds if a base seed is provided
seeds = (torch.arange(batch_size, device=device, dtype=torch.int32) + seed).int()
grid = (batch_size,)
# Power-of-2 block sizes suitable for tl.sort and modern GPUs.
# A merge buffer of 1024 (512+512) is efficient for block-level sorting.
BLOCK_SIZE_M = 512
TOP_K_BUFFER_SIZE = 512
top_k_top_p_sampling_from_probs_v128256_kernel[grid](
probs,
top_k,
top_p,
samples,
seeds,
batch_size=batch_size,
stride_probs_b=probs.stride(0),
VOCAB_SIZE=vocab_size,
BLOCK_SIZE_M=BLOCK_SIZE_M,
TOP_K_BUFFER_SIZE=TOP_K_BUFFER_SIZE,
)
# --- Output Device Management ---
return samples.to(original_device, non_blocking=True)scrolls · 246 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON