claude-opus-4-1-20250805 / tritonafd42d
claude-opus-4-1-20250805_triton_afd42d · claude-opus-4-1-20250805 · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 315 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-claude-opus-4-1-20250805-triton-afd42d?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:0625c161771b72efef988a487d1ee57bac17a0f822f66051c437ad7c8378fa00
license declaredApache-2.0
license concludedApache-2.0
authorsclaude-opus-4-1-20250805
imported2026-08-20
Kernel source
main.py315 lines
import torch
import triton
import triton.language as tl
@triton.jit
def top_k_sampling_kernel(
probs_ptr,
top_k_ptr,
samples_ptr,
seeds_ptr,
batch_size,
vocab_size: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
# Each program handles one sequence
pid = tl.program_id(0)
if pid >= batch_size:
return
# Load k for this sequence
k = tl.load(top_k_ptr + pid).to(tl.int32)
# Load random seed for this sequence
seed = tl.load(seeds_ptr + pid)
# If k is invalid, sample from full distribution
if k <= 0 or k >= vocab_size:
k = vocab_size
# We'll do multiple passes to find top-k values
# First pass: find maximum
max_val = 0.0
for block_start in range(0, vocab_size, BLOCK_SIZE):
block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = block_offsets < vocab_size
probs_offsets = pid * vocab_size + block_offsets
block_probs = tl.load(probs_ptr + probs_offsets, mask=mask, other=0.0)
block_max = tl.max(block_probs, axis=0)
max_val = tl.maximum(max_val, block_max)
# Binary search for threshold that gives us exactly k elements
# We'll find the k-th largest value
low = 0.0
high = max_val
threshold = max_val
for _ in range(20): # 20 iterations should be enough for convergence
mid = (low + high) / 2.0
# Count how many elements are >= mid
count = 0
for block_start in range(0, vocab_size, BLOCK_SIZE):
block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = block_offsets < vocab_size
probs_offsets = pid * vocab_size + block_offsets
block_probs = tl.load(probs_ptr + probs_offsets, mask=mask, other=0.0)
above_threshold = (block_probs >= mid).to(tl.int32)
count += tl.sum(above_threshold, axis=0)
if count > k:
low = mid
else:
high = mid
threshold = mid
# Now compute sum of top-k probabilities for renormalization
sum_topk = 0.0
for block_start in range(0, vocab_size, BLOCK_SIZE):
block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = block_offsets < vocab_size
probs_offsets = pid * vocab_size + block_offsets
block_probs = tl.load(probs_ptr + probs_offsets, mask=mask, other=0.0)
# Only keep probabilities >= threshold
keep_mask = block_probs >= threshold
filtered_probs = tl.where(keep_mask, block_probs, 0.0)
sum_topk += tl.sum(filtered_probs, axis=0)
# Prevent division by zero
if sum_topk <= 0.0:
sum_topk = 1.0
# Generate random number for sampling
random_offset = pid * 4 + tl.arange(0, 1)
random_val = tl.rand(seed, random_offset) * sum_topk
# Perform sampling by accumulating probabilities
cumsum = 0.0
sampled_idx = 0
found = 0
for block_start in range(0, vocab_size, BLOCK_SIZE):
block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = block_offsets < vocab_size
probs_offsets = pid * vocab_size + block_offsets
block_probs = tl.load(probs_ptr + probs_offsets, mask=mask, other=0.0)
# Only keep probabilities >= threshold
keep_mask = block_probs >= threshold
filtered_probs = tl.where(keep_mask, block_probs, 0.0)
# Compute cumulative sum for this block
# We need to process elements sequentially for cumsum
# Use a reduction approach instead
prev_cumsum = cumsum
block_cumsum = tl.cumsum(filtered_probs, axis=0) + prev_cumsum
cumsum = prev_cumsum + tl.sum(filtered_probs, axis=0)
# Check if random value falls in this block
sample_mask = (block_cumsum >= random_val) & (filtered_probs > 0) & (found == 0)
# Find the first position where sample_mask is true
# We'll use a reduction to find the minimum index where condition is true
indices_where_true = tl.where(sample_mask, block_offsets, vocab_size)
min_idx = tl.min(indices_where_true, axis=0)
if min_idx < vocab_size:
sampled_idx = min_idx
found = 1
# Fallback: if no sample was found (shouldn't happen), sample the first valid token
if found == 0:
sampled_idx = 0
# Store the sampled index
tl.store(samples_ptr + pid, sampled_idx)
@triton.jit
def top_k_sampling_kernel_simple(
probs_ptr,
top_k_ptr,
samples_ptr,
seeds_ptr,
batch_size,
vocab_size: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
"""Simplified version that's more robust"""
# Each program handles one sequence
pid = tl.program_id(0)
if pid >= batch_size:
return
# Load k for this sequence
k = tl.load(top_k_ptr + pid).to(tl.int32)
# Load random seed for this sequence
seed = tl.load(seeds_ptr + pid)
# If k is invalid, sample from full distribution
if k <= 0 or k >= vocab_size:
k = vocab_size
# Find the k-th largest value using sorting approach
# We'll use a simpler approach: find threshold iteratively
# First, find min and max values
min_val = 1.0
max_val = 0.0
for block_start in range(0, vocab_size, BLOCK_SIZE):
block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = block_offsets < vocab_size
probs_offsets = pid * vocab_size + block_offsets
block_probs = tl.load(probs_ptr + probs_offsets, mask=mask, other=0.0)
block_max = tl.max(tl.where(mask, block_probs, 0.0), axis=0)
block_min = tl.min(tl.where(mask & (block_probs > 0), block_probs, 1.0), axis=0)
max_val = tl.maximum(max_val, block_max)
min_val = tl.minimum(min_val, block_min)
# Binary search for the k-th largest value
threshold = min_val
for _ in range(30): # More iterations for better precision
mid = (min_val + max_val) / 2.0
# Count elements >= mid
count = 0
for block_start in range(0, vocab_size, BLOCK_SIZE):
block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = block_offsets < vocab_size
probs_offsets = pid * vocab_size + block_offsets
block_probs = tl.load(probs_ptr + probs_offsets, mask=mask, other=0.0)
above = (block_probs >= mid).to(tl.int32)
count += tl.sum(above, axis=0)
if count > k:
min_val = mid
else:
max_val = mid
threshold = mid
# Compute sum for renormalization
sum_topk = 0.0
for block_start in range(0, vocab_size, BLOCK_SIZE):
block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = block_offsets < vocab_size
probs_offsets = pid * vocab_size + block_offsets
block_probs = tl.load(probs_ptr + probs_offsets, mask=mask, other=0.0)
filtered = tl.where(block_probs >= threshold, block_probs, 0.0)
sum_topk += tl.sum(filtered, axis=0)
# Generate random value
random_offset = pid
rand_val = tl.rand(seed, random_offset + tl.arange(0, 1)) * sum_topk
rand_scalar = tl.sum(rand_val, axis=0) # Convert to scalar
# Sample using cumsum
cumsum = 0.0
result = vocab_size - 1 # Default to last token
for block_start in range(0, vocab_size, BLOCK_SIZE):
block_offsets = block_start + tl.arange(0, BLOCK_SIZE)
mask = block_offsets < vocab_size
probs_offsets = pid * vocab_size + block_offsets
block_probs = tl.load(probs_ptr + probs_offsets, mask=mask, other=0.0)
filtered = tl.where(block_probs >= threshold, block_probs, 0.0)
# Check each position using vectorized operations
prev_cumsum = cumsum
cumsum_vec = tl.cumsum(filtered, axis=0) + prev_cumsum
cumsum = prev_cumsum + tl.sum(filtered, axis=0)
# Find first position where cumsum >= random
above_random = cumsum_vec >= rand_scalar
valid = above_random & mask & (filtered > 0)
# Get minimum index where condition is true
indices = tl.where(valid, block_offsets, vocab_size)
min_idx = tl.min(indices, axis=0)
# Update result if we found a valid index
if min_idx < vocab_size and min_idx < result:
result = min_idx
# Store result
tl.store(samples_ptr + pid, result)
def run(probs, top_k):
"""
Top-k sampling from probability distributions.
Args:
probs: [batch_size, vocab_size] probability distributions
top_k: [batch_size] number of top tokens to consider
Returns:
samples: [batch_size] sampled token indices
"""
# Store original device
original_device = probs.device
# Move to GPU if needed
if not probs.is_cuda:
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available but GPU operation is required")
probs = probs.cuda()
if not top_k.is_cuda:
top_k = top_k.cuda() if torch.cuda.is_available() else top_k
if not top_k.is_cuda:
raise RuntimeError("CUDA is not available but GPU operation is required")
# Validate inputs
batch_size, vocab_size = probs.shape
assert vocab_size == 129280, f"Expected vocab_size=129280, got {vocab_size}"
assert top_k.shape == (batch_size,), f"top_k shape mismatch"
# Convert to required dtypes
probs = probs.to(torch.float32)
top_k = top_k.to(torch.int32)
# Allocate output
samples = torch.empty(batch_size, dtype=torch.int64, device=probs.device)
# Generate random seeds
seeds = torch.randint(0, 2**31-1, (batch_size,), dtype=torch.int32, device=probs.device)
# Choose block size - 512 works well for this vocab size
BLOCK_SIZE = 512
# Launch kernel
grid = (batch_size,)
top_k_sampling_kernel_simple[grid](
probs,
top_k,
samples,
seeds,
batch_size,
vocab_size,
BLOCK_SIZE=BLOCK_SIZE,
)
# Move result back to original device if needed
if original_device.type != 'cuda':
samples = samples.cpu()
return samplesscrolls · 315 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON