gemini-2.5-pro / triton2c9c7d
gemini-2.5-pro_triton_2c9c7d · gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 276 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-2c9c7d?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:e1a74a6add68b5a9f85b0df6980cba14d5dee87c8f9b9dcf494548007188f89c
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.
num-warps = 4
num_warps = 4shared-memory
smem_combined_packed = tl.zeros((COMBINED_SIZE,), dtype=tl.uint64, scope='shared')Kernel source
main.py276 lines
import torch
import triton
import triton.language as tl
import math
# --- Triton Kernel ---
@triton.jit
def _bitonic_sort_step(data_ptr, size, stride, merge_size, ascending):
"""
Performs one step of a bitonic sort on a 1D array in memory.
This is designed to be called iteratively to sort an array.
"""
# Each thread handles one comparison-swap operation.
# We only need size // 2 threads for this.
# However, for simplicity in Triton, we launch `size` threads and mask them.
# A more advanced implementation might use fewer threads.
idx = tl.program_id(1) * 32 + tl.arange(0, 32)
# Determine which pairs to compare
group_idx = idx // stride
inner_idx = idx % stride
# Calculate indices for comparison based on the bitonic sort network structure
i = group_idx * stride * 2 + inner_idx
j = i + stride
# Ensure we are within a merge block of size merge_size
# The direction of comparison depends on which half of the merge block we are in
is_upper_half = ((i // merge_size) % 2 == 1)
# Create a mask to avoid out-of-bounds access and redundant computations
mask = (idx < size // 2)
# Load elements to be compared
x1 = tl.load(data_ptr + i, mask=mask)
x2 = tl.load(data_ptr + j, mask=mask)
# Determine swap condition based on the bitonic sequence and final sort order
should_swap = (x1 > x2)
# Flip the swap condition based on the desired final sort order and bitonic stage
if ascending:
swap_condition = should_swap if not is_upper_half else not should_swap
else: # descending
swap_condition = should_swap if is_upper_half else not should_swap
# Perform conditional swap
swapped_x1 = tl.where(swap_condition, x2, x1)
swapped_x2 = tl.where(swap_condition, x1, x2)
# Store back the swapped elements
tl.store(data_ptr + i, swapped_x1, mask=mask)
tl.store(data_ptr + j, swapped_x2, mask=mask)
@triton.jit
def _bitonic_sort_power_of_2(data_ptr, size, ascending):
"""
Sorts a 1D tl.tensor of a power-of-2 size using a bitonic sorting network.
`data_ptr` should be a pointer to an array in shared memory.
This kernel is launched with enough threads to cover the comparisons needed.
"""
num_stages = tl.static_log2(size)
for stage in range(num_stages):
merge_size = 1 << (stage + 1)
for step in range(stage + 1):
stride = 1 << (stage - step)
# This is a conceptual call; the logic is inlined for Triton's JIT.
# In a real Triton implementation, this would be part of the main kernel loop.
# For this structure, we assume the sorting logic is called within the kernel.
# The body of `_bitonic_sort_step` would be here, or called as a utility.
# Let's assume the logic is inlined for simplicity of the demonstration.
# Inlined _bitonic_sort_step logic for one thread block:
idx = tl.arange(0, size // 2)
group_idx = idx // stride
inner_idx = idx % stride
i = group_idx * stride * 2 + inner_idx
j = i + stride
is_upper_half = ((i // merge_size) % 2 == 1)
x1 = tl.load(data_ptr + i)
x2 = tl.load(data_ptr + j)
should_swap = (x1 > x2)
if ascending:
swap_condition = should_swap if not is_upper_half else not should_swap
else:
swap_condition = should_swap if is_upper_half else not should_swap
swapped_x1 = tl.where(swap_condition, x2, x1)
swapped_x2 = tl.where(swap_condition, x1, x2)
tl.store(data_ptr + i, swapped_x1)
tl.store(data_ptr + j, swapped_x2)
tl.sync_threads()
@triton.jit
def _top_k_sampling_kernel(
probs_ptr,
top_k_ptr,
samples_ptr,
seed_tensor_ptr,
stride_probs_b,
VOCAB_SIZE: tl.constexpr,
STATIC_K: tl.constexpr,
BLOCK_V: tl.constexpr,
):
"""
Triton kernel for top-k sampling. Each program instance processes one sequence.
"""
pid_b = tl.program_id(0)
# --- Shared Memory Declaration ---
COMBINED_SIZE = tl.constexpr(STATIC_K + BLOCK_V)
smem_combined_packed = tl.zeros((COMBINED_SIZE,), dtype=tl.uint64, scope='shared')
# --- Load `k` and Seed for the current sequence ---
k = tl.load(top_k_ptr + pid_b)
seed = tl.load(seed_tensor_ptr + pid_b)
# --- Conditional execution: Top-K path vs. Full Vocab Path ---
if (k > 0) and (k < VOCAB_SIZE):
# --- Top-K Path ---
top_k_packed = tl.full([STATIC_K], 0, dtype=tl.uint64) # Start with prob=0, idx=0
v_offsets = tl.arange(0, BLOCK_V)
for v_start_idx in range(0, tl.cdiv(VOCAB_SIZE, BLOCK_V)):
v_start = v_start_idx * BLOCK_V
v_range = v_start + v_offsets
v_mask = v_range < VOCAB_SIZE
probs = tl.load(probs_ptr + pid_b * stride_probs_b + v_range, mask=v_mask, other=0.0)
indices = v_range.to(tl.uint32)
probs_uint32 = tl.view(probs, tl.uint32)
current_packed = (probs_uint32.to(tl.uint64) << 32) | indices.to(tl.uint64)
# Merge candidates in shared memory
tl.store(smem_combined_packed + tl.arange(0, STATIC_K), top_k_packed)
tl.store(smem_combined_packed + STATIC_K + v_offsets, current_packed, mask=v_mask)
tl.sync_threads()
_bitonic_sort_power_of_2(smem_combined_packed, COMBINED_SIZE, ascending=False)
top_k_packed = tl.load(smem_combined_packed + tl.arange(0, STATIC_K))
# Unpack the final top K candidates
top_k_indices = (top_k_packed & 0xFFFFFFFF).to(tl.int64)
top_k_probs = tl.view((top_k_packed >> 32).to(tl.uint32), tl.float32)
# Gumbel-Max sampling on the filtered top K items
k_arange = tl.arange(0, STATIC_K)
k_mask = k_arange < k
log_probs = tl.log(top_k_probs + 1e-9)
rand_offsets = pid_b * STATIC_K + k_arange
rand_uniform = tl.rand(seed, rand_offsets)
gumbel_noise = -tl.log(-tl.log(rand_uniform + 1e-9) + 1e-9)
gumbel_scores = tl.where(k_mask, log_probs + gumbel_noise, -float('inf'))
winner_idx_in_block = tl.argmax(gumbel_scores, axis=0)
sampled_token_id = tl.load(top_k_indices + winner_idx_in_block)
else:
# --- Full Vocab Path ---
max_gumbel_score = -float('inf')
result_index = -1
v_offsets = tl.arange(0, BLOCK_V)
for v_start_idx in range(0, tl.cdiv(VOCAB_SIZE, BLOCK_V)):
v_start = v_start_idx * BLOCK_V
v_range = v_start + v_offsets
v_mask = v_range < VOCAB_SIZE
probs = tl.load(probs_ptr + pid_b * stride_probs_b + v_range, mask=v_mask, other=0.0)
log_probs = tl.log(probs + 1e-9)
rand_offsets = pid_b * VOCAB_SIZE + v_range
rand_uniform = tl.rand(seed, rand_offsets)
gumbel_noise = -tl.log(-tl.log(rand_uniform + 1e-9) + 1e-9)
gumbel_scores = tl.where(v_mask, log_probs + gumbel_noise, -float('inf'))
block_max_score = tl.max(gumbel_scores, axis=0)
update_mask = block_max_score > max_gumbel_score
max_gumbel_score = tl.where(update_mask, block_max_score, max_gumbel_score)
block_max_idx = tl.argmax(gumbel_scores, axis=0)
block_winner_vocab_idx = (v_start + block_max_idx)
result_index = tl.where(update_mask, block_winner_vocab_idx, result_index)
sampled_token_id = result_index.to(tl.int64)
tl.store(samples_ptr + pid_b, sampled_token_id)
def top_k_sampling_from_probs_v129280(probs: torch.Tensor, top_k: torch.Tensor) -> torch.Tensor:
"""
Performs top-k sampling from probability distributions using a Triton kernel.
"""
if not torch.cuda.is_available():
raise RuntimeError("This kernel requires a CUDA-enabled GPU.")
if probs.dim() != 2 or top_k.dim() != 1 or probs.shape[0] != top_k.shape[0]:
raise ValueError("Invalid shapes. probs must be [batch, vocab], top_k must be [batch].")
batch_size, vocab_size = probs.shape
assert vocab_size == 129280, "This kernel is specialized for vocab_size=129280"
# Define kernel constants.
# Note: Using larger STATIC_K might require more shared memory and register spills,
# but handles larger k values more efficiently within the fast path.
# For bitonic sort, (STATIC_K + BLOCK_V) must be a power of 2.
STATIC_K = 64
BLOCK_V = 64
combined_size = STATIC_K + BLOCK_V
if (combined_size & (combined_size - 1) != 0) or combined_size == 0:
raise ValueError(f"STATIC_K ({STATIC_K}) + BLOCK_V ({BLOCK_V}) must be a power of two for the bitonic sort.")
original_device = probs.device
device = torch.device("cuda")
# Move data to GPU
probs_gpu = probs.to(device=device, dtype=torch.float32, non_blocking=True)
top_k_gpu = top_k.to(device=device, dtype=torch.int32, non_blocking=True)
# Allocate output and seed tensors
samples = torch.empty(batch_size, dtype=torch.int64, device=device)
seed_tensor = torch.randint(0, 2**32 - 1, (batch_size,), dtype=torch.int64, device=device)
grid = (batch_size,)
# We use one warp per program instance. More complex kernels might need more.
# The bitonic sort implementation implicitly uses all threads in the block.
# A single warp (32 threads) is sufficient for vector loads/stores.
# However, the bitonic sort is most efficient when using more threads.
# Let's use 4 warps to provide enough parallelism for the sort.
num_warps = 4
_top_k_sampling_kernel[grid](
probs_gpu,
top_k_gpu,
samples,
seed_tensor,
stride_probs_b=probs_gpu.stride(0),
VOCAB_SIZE=vocab_size,
STATIC_K=STATIC_K,
BLOCK_V=BLOCK_V,
num_warps=num_warps,
)
# Move result back to the original device
return samples.to(device=original_device, non_blocking=True)
def run(*args, **kwargs):
"""
Public entry point for the kernel.
Handles device management and positional/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:
try:
probs = kwargs['probs']
top_k = kwargs['top_k']
except KeyError as e:
raise KeyError(f"Missing required keyword argument: {e}")
else:
raise ValueError("No arguments provided. Please provide 'probs' and 'top_k'.")
return top_k_sampling_from_probs_v129280(probs, top_k)scrolls · 276 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON